64 Commits
Author SHA1 Message Date
mengli.cml 89cae4def0 add gradient norm tensorboard visualizatoin for deepspeed mode 2025-01-14 18:13:11 +08:00
mengli.cml a35d7bbf0e fix bug 2025-01-09 10:29:23 +08:00
mengli.cml bfce9c2583 support for parallel inference using xfuser 2025-01-06 20:45:10 +08:00
Bubbliiiing 45c298af71 Update Readme && Update lora cache && Lora merge with CUDA (#173)
* Update Readme

* Fix bug in training

* Update Readme

* Update Readme && Update lora cache

* Update Readme

* Update error info in sampler

* Fix Memory bug in 12b
2024-12-26 15:41:21 +08:00
2eb809a0c8 Add lora cache mode in comfyui && Use cuda in merge_lora func for speeding up. (#172)
* fix lora precision problem during multiple merge unmerge, keep the original weights on cpu

* add comments

* update

* Update Lora cache

---------

Co-authored-by: mengli.cml <mengli.cml@alibaba-inc.com>
Co-authored-by: bubbliiiing <3323290568@qq.com>
2024-12-24 10:02:57 +08:00
hkz 58b27cee1c fix caption rewrite (#169) 2024-12-12 18:21:58 +08:00
hkz 485e771acf Fix train reward lora and video caption (#165) 2024-12-10 14:23:33 +08:00
Bubbliiiing ccbdce4492 Update Readme (#161)
* Update Readme

* Fix bug in training

* Update Readme
2024-12-02 11:27:57 +08:00
hkzandbubbliiiing 17c1f02ba8 Add reward LoRA training (#160)
* Add reward lora

* Fix bug in lora utils && Update Readme

---------

Co-authored-by: bubbliiiing <3323290568@qq.com>
2024-11-28 13:59:53 +08:00
Bubbliiiing f419bf850b Update EasyAnimateV5-7B and reduce vram in vae decoding (#158) 2024-11-22 14:20:29 +08:00
Bubbliiiing 6d15f2799f Update readme (#155) 2024-11-20 16:03:55 +08:00
Bubbliiiing 94f420c257 Low ram support. (#154)
* Fix bug in install comfyui

* support model load in 30GB RAM

* Update requirement

* Update requirement

* Update VAE outlier_penalty_loss

* Fix bug in ui

* Fix bug in ui

* Fix bug in ui
2024-11-20 15:56:28 +08:00
Bubbliiiing 9d8fed9cec Fix bug in install comfyui (#148) 2024-11-13 19:49:18 +08:00
hkz 30682ba3c5 Update video caption (#129) 2024-11-13 17:43:21 +08:00
Bubbliiiing 5dde5c9a7e Update Reference (#143)
* Fix bug in comfyui t2v

* Fix bug in textbox

* Update Reference
2024-11-12 12:57:30 +08:00
Bubbliiiing 59bb7f0169 Bug fix in TextBox (#141)
* Fix bug in comfyui t2v

* Fix bug in textbox
2024-11-12 10:59:18 +08:00
Bubbliiiing 2f76bace4d Fix bug in comfyui t2v (#140) 2024-11-12 10:28:13 +08:00
yunkchen 3e3e4756ce Update README.md (#139) 2024-11-12 10:25:00 +08:00
Ikko Eltociear Ashimine 95f11a5403 docs: add Japanese README file (#138)
I created Japanese translated README.
2024-11-12 10:19:27 +08:00
Bubbliiiing 9da6b79ece Update Prompt && Update Readme && Fix control bug (#131)
* Update Readme

* Update prompt tips

* Fix bug in ui

* Fix control bug and Update readme

* Update prompt again
2024-11-11 13:48:53 +08:00
Bubbliiiing 44539fe77b Update Readme (#130) 2024-11-08 22:38:17 +08:00
Bubbliiiing 62de94e2f1 Update V5 (#128)
Update V5
2024-11-08 19:46:50 +08:00
yunkchen abcf42bce7 Update requirements.txt (#124)
Ensure gradio version.
2024-10-18 17:04:56 +08:00
Bubbliiiing d3b8bbbd14 Fix bug in v3 training (#98) 2024-08-22 16:03:14 +08:00
hkz 5ea1bf2450 Update Video Caption (#93)
* Add VILA1.5 in video caption

* update VILA1.5 in video_caption

* add get_video_path_list back & update the accelerate version & auto load the vila quant_path & download models in the main process

* add exception handling & add force_ascii=False

* fix the NCCL timeout & use logger

* update utils

* Update video splitting

* Update pre-filtering

* Update video caption

* Add caption_rewrite.py

* Add VideoCLIPXL

* Add beautiful prompt demo

* fix test

* Update Dockerfile.ds

* update VideoCLIPXL & add filter_meta_train.py

* update stage3

* update README.md

* update README_zh-CN.md

* update README

* update requirements

* fix stage_3_video_recaptioning.sh

* update Beautiful Prompt

* update README

* update vila_video_recaptioning.py & fix the empty gather result
2024-08-21 16:40:54 +08:00
Bubbliiiing f2f0cbccc9 Bug fix/encode prompt (#96)
* update readme

* fix bug in encode_prompt
2024-08-19 16:53:05 +08:00
Bubbliiiing aa3d2afb34 update readme (#95) 2024-08-19 12:02:43 +08:00
Bubbliiiing c7c2fc1cd0 update readme (#94) 2024-08-19 11:26:41 +08:00
e4f1a4fe97 Update V4 version (#92)
* update train_lora && update deepspeed && update training with max token length

* fix bug in train.py

* fix bug in training_with_video_token_length

* update v2v && update v2v api

* add rope2d embedding precomputation; move text encoder to dataloader to reduce gpu memory consumpution

* add cuda multi-stream to speedup vae encode

* update new vae && new comfyui

* fix some bug in training code

* Add lcm lora (#89)

Co-authored-by: xuanyuan.lb <xuanyuan.lb@alibaba-inc.com>

* Update Training Code and fix bug in low vram mode

* fix bug in low vram mode

* update report

* update cfg

* actual text clip

---------

Co-authored-by: mengli.cml <mengli.cml@alibaba-inc.com>
Co-authored-by: liubo0902 <38622806+liubo0902@users.noreply.github.com>
Co-authored-by: xuanyuan.lb <xuanyuan.lb@alibaba-inc.com>
2024-08-19 11:22:17 +08:00
Bubbliiiing b54412ceb0 support float16 in comfyui (#75) 2024-07-24 14:01:35 +08:00
bubbliiiing 039d67acf1 support float16 && add reference && fix bug in training 2024-07-24 13:19:03 +08:00
bubbliiiing 1a5bab2234 fix bug in no inpaint model 2024-07-18 21:09:10 +08:00
bubbliiiing 2eef3f9f78 update v4 2024-07-18 16:33:56 +08:00
Bubbliiiing 616a35425d Fix the issue where Lora training cannot load state, fix the issue where training cannot eval (#66)
* rename the comfyui files and new readme

* fix bug in eval and load lora state dict
2024-07-18 15:47:51 +08:00
Bubbliiiing 8b7722463e rename the comfyui files and new readme (#53) 2024-07-13 14:57:13 +08:00
Bubbliiiing 883c0a20a8 update readme (#52) 2024-07-12 18:03:48 +08:00
Bubbliiiing e6ec6f0e04 Update init and install in comfyui (#51)
Update init and install in comfyui
2024-07-12 17:58:12 +08:00
Bubbliiiing fbbfc818ea Update pipeline (#50) 2024-07-12 17:47:47 +08:00
yunkchen 3cc15c456a Merge pull request #49 from aigc-apps/comfyui
Support for comfyui
2024-07-12 16:54:06 +08:00
bubbliiiing deaa46ff42 update readme for comfyui 2024-07-12 16:31:50 +08:00
bubbliiiing 8fd346a9af Merge branch 'main' into comfyui 2024-07-12 16:26:34 +08:00
bubbliiiing c1d7e503d3 update comfyui 2024-07-12 16:24:46 +08:00
yunkchenandbubbliiiiing 047712c8bd Add Discord (#48)
Add Discord
---------

Co-authored-by: bubbliiiiing <3323290568@qq.com>
2024-07-12 11:41:24 +08:00
Bubbliiiing 6ed8619ad4 Rename the files (#45)
* add update readme

* update demo show and prompt

* rename the train code
2024-07-10 17:46:02 +08:00
bubbliiiing 4f6bd0f6a6 rename the train code 2024-07-10 17:42:12 +08:00
yunkchen 19b38f674e Merge pull request #44 from wangqiang9/main
Added exception handling for video read bucket_sampler.py
2024-07-10 16:35:09 +08:00
Wang Qiang d57c779217 Added exception handling for video read bucket_sampler.py 2024-07-10 15:42:30 +08:00
Bubbliiiing f1e013619c Update Readme and Gallery (#42)
* add update readme

* update demo show and prompt
2024-07-06 13:10:13 +08:00
bubbliiiing a396e2cee4 update demo show and prompt 2024-07-06 13:07:37 +08:00
bubbliiiiing 9e55766f06 add update readme 2024-07-06 12:07:30 +08:00
Bubbliiiingandchenyunkuo.cyk f9eeabe231 Updated to v3 version, supports image generated videos, with a maximum support of 960x960x144 video generation. (#40)
* update v3

* Fix frame start_idx bug, see issue #41.

* update readme and fix bug in training

* update requirements

* update ui

* fix bug in inpaint

* update new ui

* fix bug in auto resize

* fix bug in auto resize

* fix bug in modelscope and eas

* update low gpu memory mode

---------

Co-authored-by: chenyunkuo.cyk <chenyunkuo.cyk@alibaba-inc.com>
2024-07-05 20:27:24 +08:00
zouxinyi0625 e4c824a1c5 add dj news (#30) 2024-06-17 15:55:13 +08:00
bubbliiiing 59bd5de1ad update inpaint model 2024-06-08 09:38:48 +08:00
bubbliiiing 60d205111b update inpaint model and new ui in demo 2024-06-07 10:48:13 +08:00
Bubbliiiing d196f64d93 Add huggingface link and new UI (#22) 2024-06-04 21:33:15 +08:00
zouxinyi0625 dc9305b0fe update ui chinese (#20)
* update ui chinese
2024-06-04 16:24:56 +08:00
Bubbliiiing 58d793659d Update 768x768 Link and new gallery (#15)
update readme and model link
2024-06-04 10:47:57 +08:00
hkz 96ebec14c1 Merge pull request #10 from aigc-apps/fix_autogptq_sglang
fix the conflict between autogptq and sglang
2024-05-31 15:30:30 +08:00
hkunzhe 9e2fba7d40 fix the conflict between autogptq and sglang 2024-05-31 14:27:56 +08:00
zouxinyi0625 480e52c796 update readme (#8)
update arxiv link
2024-05-31 11:15:09 +08:00
Bubbliiiingandzouxinyi0625 fb0916f4da Update EasyAnimateV2 (#5)
* update EasyAnimateV2

* update datasets loader

* update datasets loader

* update fast api

* complete data preprocess pipeline.

* update lots of readme

* update readme bans

* fix bug in validation while training

* provide example for video cut

* update text box

* add arxiv

* delete IDDPM

* update gallery

* update arxiv

* update readme

* update readme

* link fix

* update vae readme

---------

Co-authored-by: zouxinyi0625 <zouxinyi.zxy@alibaba-inc.com>
2024-05-31 10:33:14 +08:00
Bubbliiiing 362521e5da Merge pull request #6 from aigc-apps/add_dataset_pipeline
add the dataset preprocessing pipeline
2024-05-28 20:33:32 +08:00
hkunzhe a8cfe4ec7e fix test 2024-05-28 12:00:49 +08:00
hkunzhe 6efc7de3f2 add the dataset preprocessing pipeline 2024-05-27 22:26:03 +08:00
203 changed files with 48173 additions and 9050 deletions
+3
View File
@@ -2,6 +2,9 @@
models*
output*
samples*
datasets*
asset*
_*
__pycache__/
*.py[cod]
*$py.class
+30 -10
View File
@@ -1,4 +1,4 @@
FROM nvidia/cuda:11.8.0-devel-ubuntu22.04
FROM nvidia/cuda:12.1.0-cudnn8-devel-ubuntu22.04
ENV DEBIAN_FRONTEND noninteractive
RUN rm -r /etc/apt/sources.list.d/
@@ -6,26 +6,46 @@ RUN rm -r /etc/apt/sources.list.d/
RUN apt-get update -y && apt-get install -y \
libgl1 libglib2.0-0 google-perftools \
sudo wget git git-lfs vim tig pkg-config libcairo2-dev \
telnet curl net-tools iputils-ping wget jq \
python3-pip python-is-python3 python3.10-venv tzdata lsof && \
rm -rf /var/lib/apt/lists/*
aria2 telnet curl net-tools iputils-ping jq \
python3-pip python-is-python3 python3.10-venv tzdata lsof zip tmux
RUN apt-get update && \
apt-get install -y software-properties-common && \
add-apt-repository ppa:ubuntuhandbook1/ffmpeg6 && \
apt-get update && \
apt-get install -y ffmpeg
RUN pip3 install --upgrade pip -i https://mirrors.aliyun.com/pypi/simple/
# add all extensions
RUN apt-get update -y && apt-get install -y zip && \
rm -rf /var/lib/apt/lists/*
RUN pip install wandb tqdm GitPython==3.1.32 Pillow==9.5.0 setuptools --upgrade -i https://mirrors.aliyun.com/pypi/simple/
# reinstall torch to keep compatible with xformers
RUN pip uninstall -qy torch torchvision && \
pip install torch==2.2.0 torchvision==0.17.0 torchaudio==2.2.0 --index-url https://download.pytorch.org/whl/cu118
RUN pip uninstall -qy xfromers && pip install xformers==0.0.24 --index-url https://download.pytorch.org/whl/cu118
RUN pip install torch==2.4.0 torchvision==0.19.0 torchaudio==2.4.0 --index-url https://download.pytorch.org/whl/cu118
RUN pip install xformers==0.0.27.post2 --index-url https://download.pytorch.org/whl/cu118
# install vllm (video-caption)
RUN pip install vllm==0.6.3
# install requirements (video-caption)
WORKDIR /root/
COPY easyanimate/video_caption/requirements.txt /root/requirements-video_caption.txt
RUN pip install -r /root/requirements-video_caption.txt
RUN rm /root/requirements-video_caption.txt
RUN pip install -U http://eas-data.oss-cn-shanghai.aliyuncs.com/sdk/allspark-0.15-py2.py3-none-any.whl
RUN pip install -e git+https://github.com/CompVis/taming-transformers.git@master#egg=taming-transformers
RUN pip install came-pytorch deepspeed pytorch_lightning==1.9.4 func_timeout -i https://mirrors.aliyun.com/pypi/simple/
# install requirements
RUN pip install bitsandbytes mamba-ssm causal-conv1d>=1.4.0 -i https://mirrors.aliyun.com/pypi/simple/
RUN pip install ipykernel -i https://mirrors.aliyun.com/pypi/simple/
COPY ./requirements.txt /root/requirements.txt
RUN pip install -r /root/requirements.txt -i https://mirrors.aliyun.com/pypi/simple/
RUN rm -rf /root/requirements.txt
# install package patches (video-caption)
COPY easyanimate/video_caption/package_patches/easyocr_detection_patched.py /usr/local/lib/python3.10/dist-packages/easyocr/detection.py
COPY easyanimate/video_caption/package_patches/vila_siglip_encoder_patched.py /usr/local/lib/python3.10/dist-packages/llava/model/multimodal_encoder/siglip_encoder.py
ENV PYTHONUNBUFFERED 1
ENV NVIDIA_DISABLE_REQUIRE 1
+357 -300
View File
@@ -1,107 +1,67 @@
# 📷 EasyAnimate | Your Animation Generator.
😊 EasyAnimate is a repo for generating long videos and training transformer based diffusion generators.
# 📷 EasyAnimate | An End-to-End Solution for High-Resolution and Long Video Generation
😊 EasyAnimate is an end-to-end solution for generating high-resolution and long videos. We can train transformer based diffusion generators, train VAEs for processing long videos, and preprocess metadata.
😊 Based on Sora like structure and DIT, we use transformer as a diffuser for video generation. In order to ensure good expansibility, we built easyanimate based on motion module. In the future, we will try more training programs to improve the effect.
😊 We use DIT and transformer as a diffuser for video and image generation.
😊 Welcome!
English | [简体中文](./README_zh-CN.md)
[![Arxiv Page](https://img.shields.io/badge/Arxiv-Page-red)](https://arxiv.org/abs/2405.18991)
[![Project Page](https://img.shields.io/badge/Project-Website-green)](https://easyanimate.github.io/)
[![Modelscope Studio](https://img.shields.io/badge/Modelscope-Studio-blue)](https://modelscope.cn/studios/PAI/EasyAnimate/summary)
[![Hugging Face Spaces](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Spaces-yellow)](https://huggingface.co/spaces/alibaba-pai/EasyAnimate)
[![Discord Page](https://img.shields.io/badge/Discord-Page-blue)](https://discord.gg/UzkpB4Bn)
English | [简体中文](./README_zh-CN.md) | [日本語](./README_ja-JP.md)
# Table of Contents
- [Table of Contents](#table-of-contents)
- [Introduction](#introduction)
- [TODO List](#todo-list)
- [Model zoo](#model-zoo)
- [1、Motion Weights](#1motion-weights)
- [2、Other Weights](#2other-weights)
- [Quick Start](#quick-start)
- [1. Cloud usage: AliyunDSW/Docker](#1-cloud-usage-aliyundswdocker)
- [2. Local install: Environment Check/Downloading/Installation](#2-local-install-environment-checkdownloadinginstallation)
- [Video Result](#video-result)
- [How to use](#how-to-use)
- [1. Inference](#1-inference)
- [2. Model Training](#2-model-training)
- [Algorithm Detailed](#algorithm-detailed)
- [Model zoo](#model-zoo)
- [TODO List](#todo-list)
- [Contact Us](#contact-us)
- [Reference](#reference)
- [License](#license)
# Introduction
EasyAnimate is a pipeline based on the transformer architecture that can be used to generate AI animations, train baseline models and Lora models for the Diffusion Transformer. We support making predictions directly from the pre-trained EasyAnimate model to generate videos of about different resolutions, 6 seconds with 12 fps (40 ~ 80 frames, in the future, we will support longer videos). Users are also supported to train their own baseline models and Lora models to perform certain style transformations.
EasyAnimate is a pipeline based on the transformer architecture, designed for generating AI images and videos, and for training baseline models and Lora models for Diffusion Transformer. We support direct prediction from pre-trained EasyAnimate models, allowing for the generation of videos with various resolutions, approximately 6 seconds in length, at 8fps (EasyAnimateV5, 1 to 49 frames). Additionally, users can train their own baseline and Lora models for specific style transformations.
We will support quick pull-ups from different platforms, refer to [Quick Start](#quick-start).
What's New:
- Add Code for [video-caption](./easyanimate/video_caption/). [ 2024.04.17 ]
- Create Code! Support for Windows and Linux Now. [ 2024.04.12 ]
**New Features:**
- Use reward backpropagation to train Lora and optimize the video, aligning it better with human preferences, detailes in [here](scripts/README_TRAIN_REWARD.md). EasyAnimateV5-7b is released now. [2024.11.27]
- **Updated to v5**, supporting video generation up to 1024x1024, 49 frames, 6s, 8fps, with expanded model scale to 12B, incorporating the MMDIT structure, and enabling control models with diverse inputs; supports bilingual predictions in Chinese and English. [2024.11.08]
- **Updated to v4**, allowing for video generation up to 1024x1024, 144 frames, 6s, 24fps; supports video generation from text, image, and video, with a single model handling resolutions from 512 to 1280; bilingual predictions in Chinese and English enabled. [2024.08.15]
- **Updated to v3**, supporting video generation up to 960x960, 144 frames, 6s, 24fps, from text and image. [2024.07.01]
- **ModelScope-Sora “Data Director” Creative Race** — The third Data-Juicer Big Model Data Challenge is now officially launched! Utilizing EasyAnimate as the base model, it explores the impact of data processing on model training. Visit the [competition website](https://tianchi.aliyun.com/competition/entrance/532219) for details. [2024.06.17]
- **Updated to v2**, supporting video generation up to 768x768, 144 frames, 6s, 24fps. [2024.05.26]
- **Code Created!** Now supporting Windows and Linux. [2024.04.12]
These are our generated results:
Function:
- [Data Preprocessing](#data-preprocess)
- [Train VAE](#vae-train)
- [Train DiT](#dit-train)
- [Video Generation](#video-gen)
Our UI interface is as follows:
![ui](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/ui.png)
# TODO List
- Support model with larger resolution.
- Support model with magvit.
- Support video inpaint model.
# Model zoo
### 1、Motion Weights
| Name | Type | Storage Space | Url | Description |
|--|--|--|--|--|
| easyanimate_v1_mm.safetensors | Motion Module | 4.1GB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Motion_Module/easyanimate_v1_mm.safetensors) | Training with 80 frames and fps 12 |
### 2、Other Weights
| Name | Type | Storage Space | Url | Description |
|--|--|--|--|--|
| PixArt-XL-2-512x512.tar | Pixart | 11.4GB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Diffusion_Transformer/PixArt-XL-2-512x512.tar)| Pixart-Alpha official weights |
| easyanimate_portrait.safetensors | Checkpoint of Pixart | 2.3GB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait.safetensors) | Training with internal portrait datasets |
| easyanimate_portrait_lora.safetensors | Lora of Pixart | 654.0MB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait_lora.safetensors)| Training with internal portrait datasets |
# Result Gallery
When generating landscape animations, the sampler recommends using DPM++and Euler A. When generating portrait animations, the sampler recommends using Euler A and Euler.
Sometimes Github cannot display large GIFs properly. You can download GIFs locally to view them.
Work with origin transformer weights.
| Base Models | Sampler | Seed | Resolution (h x w x f) | Prompt | GenerationResult | Download |
| ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ |
| PixArt | DPM++ | 43 | 512x512x80 | A soaring drone footage captures the majestic beauty of a coastal cliff, its red and yellow stratified rock faces rich in color and against the vibrant turquoise of the sea. Seabirds can be seen taking flight around the cliff\'s precipices. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/1-cliff.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/1-cliff.gif) |
| PixArt | DPM++ | 43 | 448x640x80 | The video captures the majestic beauty of a waterfall cascading down a cliff into a serene lake. The waterfall, with its powerful flow, is the central focus of the video. The surrounding landscape is lush and green, with trees and foliage adding to the natural beauty of the scene. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/2-waterfall.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/2-waterfall.gif) |
| PixArt | DPM++ | 43 | 704x384x80 | A vibrant scene of a snowy mountain landscape. The sky is filled with a multitude of colorful hot air balloons, each floating at different heights, creating a dynamic and lively atmosphere. The balloons are scattered across the sky, some closer to the viewer, others further away, adding depth to the scene. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/3-snowy.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/3-snowy.gif) |
| PixArt | DPM++ | 43 | 448x640x64 | The vibrant beauty of a sunflower field. The sunflowers, with their bright yellow petals and dark brown centers, are in full bloom, creating a stunning contrast against the green leaves and stems. The sunflowers are arranged in neat rows, creating a sense of order and symmetry. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/4-sunflower.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/4-sunflower.gif) |
| PixArt | DPM++ | 43 | 384x704x48 | A tranquil Vermont autumn, with leaves in vibrant colors of orange and red fluttering down a mountain stream. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/5-autumn.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/5-autumn.gif) |
| PixArt | DPM++ | 43 | 704x384x48 | A vibrant underwater scene. A group of blue fish, with yellow fins, are swimming around a coral reef. The coral reef is a mix of brown and green, providing a natural habitat for the fish. The water is a deep blue, indicating a depth of around 30 feet. The fish are swimming in a circular pattern around the coral reef, indicating a sense of motion and activity. The overall scene is a beautiful representation of marine life. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/6-underwater.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/6-underwater.gif) |
| PixArt | DPM++ | 43 | 576x448x48 | Pacific coast, carmel by the blue sea ocean and peaceful waves | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/7-coast.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/7-coast.gif) |
| PixArt | DPM++ | 43 | 576x448x80 | A snowy forest landscape with a dirt road running through it. The road is flanked by trees covered in snow, and the ground is also covered in snow. The sun is shining, creating a bright and serene atmosphere. The road appears to be empty, and there are no people or animals visible in the video. The style of the video is a natural landscape shot, with a focus on the beauty of the snowy forest and the peacefulness of the road. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/8-forest.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/8-forest.gif) |
| PixArt | DPM++ | 43 | 640x448x64 | The dynamic movement of tall, wispy grasses swaying in the wind. The sky above is filled with clouds, creating a dramatic backdrop. The sunlight pierces through the clouds, casting a warm glow on the scene. The grasses are a mix of green and brown, indicating a change in seasons. The overall style of the video is naturalistic, capturing the beauty of the landscape in a realistic manner. The focus is on the grasses and their movement, with the sky serving as a secondary element. The video does not contain any human or animal elements. |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/9-grasses.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/9-grasses.gif) |
| PixArt | DPM++ | 43 | 704x384x80 | A serene night scene in a forested area. The first frame shows a tranquil lake reflecting the star-filled sky above. The second frame reveals a beautiful sunset, casting a warm glow over the landscape. The third frame showcases the night sky, filled with stars and a vibrant Milky Way galaxy. The video is a time-lapse, capturing the transition from day to night, with the lake and forest serving as a constant backdrop. The style of the video is naturalistic, emphasizing the beauty of the night sky and the peacefulness of the forest. |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/10-night.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/10-night.gif) |
| PixArt | DPM++ | 43 | 640x448x80 | Sunset over the sea. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/11-sunset.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/11-sunset.gif) |
Work with Portrait transformer weights.
| Base Models | Sampler | Seed | Resolution (h x w x f) | Prompt | GenerationResult | Download |
| ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ |
| Portrait | Euler A | 43 | 448x576x80 | 1girl, 3d, black hair, brown eyes, earrings, grey background, jewelry, lips, long hair, looking at viewer, photo \\(medium\\), realistic, red lips, solo | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/1-check.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/1-check.gif) |
| Portrait | Euler A | 43 | 448x576x80 | 1girl, bare shoulders, blurry, brown eyes, dirty, dirty face, freckles, lips, long hair, looking at viewer, realistic, sleeveless, solo, upper body |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/2-check.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/2-check.gif) |
| Portrait | Euler A | 43 | 512x512x64 | 1girl, black hair, brown eyes, earrings, grey background, jewelry, lips, looking at viewer, mole, mole under eye, neck tattoo, nose, ponytail, realistic, shirt, simple background, solo, tattoo |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/3-check.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/3-check.gif) |
| Portrait | Euler A | 43 | 576x448x64 | 1girl, black hair, lips, looking at viewer, mole, mole under eye, mole under mouth, realistic, solo |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/5-check.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/5-check.gif) |
Work with Portrait transformer Lora.
| Base Models | Sampler | Seed | Resolution (h x w x f) | Prompt | GenerationResult | Download |
| ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ |
| Pixart + Lora | Euler A | 43 | 512x512x64 | 1girl, 3d, black hair, brown eyes, earrings, grey background, jewelry, lips, long hair, looking at viewer, photo \\(medium\\), realistic, red lips, solo |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/1-lora.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/1-lora.gif) |
| Pixart + Lora | Euler A | 43 | 512x512x64 | 1girl, bare shoulders, blurry, brown eyes, dirty, dirty face, freckles, lips, long hair, looking at viewer, mole, mole on breast, mole on neck, mole under eye, mole under mouth, realistic, sleeveless, solo, upper body |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/2-lora.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/2-lora.gif) |
| Pixart + Lora | Euler A | 43 | 512x512x64 | 1girl, black hair, lips, looking at viewer, mole, mole under eye, mole under mouth, realistic, solo |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/5-lora.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/5-lora.gif) |
| Pixart + Lora | Euler A | 43 | 512x512x80 | 1girl, bare shoulders, blurry, blurry background, blurry foreground, bokeh, brown eyes, christmas tree, closed mouth, collarbone, depth of field, earrings, jewelry, lips, long hair, looking at viewer, photo \\(medium\\), realistic, smile, solo |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/8-lora.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/8-lora.gif) |
![ui](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/ui_v3.jpg)
# Quick Start
### 1. Cloud usage: AliyunDSW/Docker
#### a. From AliyunDSW
Stay tuned.
DSW has free GPU time, which can be applied once by a user and is valid for 3 months after applying.
#### b. From docker
Aliyun provide free GPU time in [Freetier](https://free.aliyun.com/?product=9602825&crowd=enterprise&spm=5176.28055625.J_5831864660.1.e939154aRgha4e&scm=20140722.M_9974135.P_110.MO_1806-ID_9974135-MID_9974135-CID_30683-ST_8512-V_1), get it and use in Aliyun PAI-DSW to start EasyAnimate within 5min!
[![DSW Notebook](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/dsw.png)](https://gallery.pai-ml.com/#/preview/deepLearning/cv/easyanimate_v5)
#### b. From ComfyUI
Our ComfyUI is as follows, please refer to [ComfyUI README](comfyui/README.md) for details.
![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v3/comfyui_i2v.jpg)
#### c. From docker
If you are using docker, please make sure that the graphics card driver and CUDA environment have been installed correctly in your machine.
Then execute the following commands in this way:
@@ -123,319 +83,416 @@ mkdir models/Diffusion_Transformer
mkdir models/Motion_Module
mkdir models/Personalized_Model
wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Motion_Module/easyanimate_v1_mm.safetensors -O models/Motion_Module/easyanimate_v1_mm.safetensors
wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait.safetensors -O models/Personalized_Model/easyanimate_portrait.safetensors
wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait_lora.safetensors -O models/Personalized_Model/easyanimate_portrait_lora.safetensors
wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Diffusion_Transformer/PixArt-XL-2-512x512.tar -O models/Diffusion_Transformer/PixArt-XL-2-512x512.tar
cd models/Diffusion_Transformer/
tar -xvf PixArt-XL-2-512x512.tar
cd ../../
# Please use the hugginface link or modelscope link to download the EasyAnimateV5 model.
# I2V models
# https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-InP
# https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-InP
# T2V models
# https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh
# https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh
```
### 2. Local install: Environment Check/Downloading/Installation
#### a. Environment Check
We have verified EasyAnimate execution on the following environment:
The detailed of Windows:
- OS: Windows 10
- python: python3.10 & python3.11
- pytorch: torch2.2.0
- CUDA: 11.8 & 12.1
- CUDNN: 8+
- GPU: Nvidia-3060 12G
The detailed of Linux:
- OS: Ubuntu 20.04, CentOS
- python: py3.10 & py3.11
- python: python3.10 & python3.11
- pytorch: torch2.2.0
- CUDA: 11.8
- CUDA: 11.8 & 12.1
- CUDNN: 8+
- GPU: Nvidia-A10 24G & Nvidia-A100 40G & Nvidia-A100 80G
- GPU:Nvidia-V100 16G & Nvidia-A10 24G & Nvidia-A100 40G & Nvidia-A100 80G
We need about 60GB available on disk (for saving weights), please check!
#### b. Weights
We'd better place the weights along the specified path:
The video size for EasyAnimateV5-12B can be generated by different GPU Memory, including:
| GPU memory |384x672x25|384x672x49|576x1008x25|576x1008x49|768x1344x25|768x1344x49|
|----------|----------|----------|----------|----------|----------|----------|
| 16GB | 🧡 | 🧡 | ❌ | ❌ | ❌ | ❌ |
| 24GB | 🧡 | 🧡 | 🧡 | 🧡 | 🧡 | ❌ |
| 40GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
| 80GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
The video size for EasyAnimateV5-7B can be generated by different GPU Memory, including:
| GPU memory |384x672x25|384x672x49|576x1008x25|576x1008x49|768x1344x25|768x1344x49|
|----------|----------|----------|----------|----------|----------|----------|
| 16GB | 🧡 | 🧡 | ❌ | ❌ | ❌ | ❌ |
| 24GB | ✅ | ✅ | ✅ | 🧡 | 🧡 | ❌ |
| 40GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
| 80GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
✅ indicates it can run under "model_cpu_offload", 🧡 represents it can run under "model_cpu_offload_and_qfloat8", ⭕️ indicates it can run under "sequential_cpu_offload", ❌ means it can't run. Please note that running with sequential_cpu_offload will be slower.
Some GPUs that do not support torch.bfloat16, such as 2080ti and V100, require changing the weight_dtype in app.py and predict files to torch.float16 in order to run.
The generation time for EasyAnimateV5-12B using different GPUs over 25 steps is as follows:
| GPU | 384x672x25 | 384x672x49 | 576x1008x25 | 576x1008x49 | 768x1344x25 | 768x1344x49 |
|-----------|------------------|------------------|------------------|------------------|------------------|-----------------|
| A10 24GB | ~120s (4.8s/it) | ~240s (9.6s/it) | ~320s (12.7s/it) | ~750s (29.8s/it) | ❌ | ❌ |
| A100 80GB | ~45s (1.75s/it) | ~90s (3.7s/it) | ~120s (4.7s/it) | ~300s (11.4s/it) | ~265s (10.6s/it) | ~710s (28.3s/it) |
(⭕️) indicates it can run with low_gpu_memory_mode=True, but at a slower speed, and ❌ means it can't run.
<details>
<summary>(Obsolete) EasyAnimateV3:</summary>
The video size for EasyAnimateV3 can be generated by different GPU Memory, including:
| GPU memory | 384x672x25 | 384x672x144 | 576x1008x72 | 576x1008x144 | 720x1280x72 | 720x1280x144 |
|------------|------------|-------------|-------------|--------------|-------------|--------------|
| 12GB | ⭕️ | ⭕️ | ⭕️ | ⭕️ | ❌ | ❌ |
| 16GB | ✅ | ✅ | ⭕️ | ⭕️ | ⭕️ | ❌ |
| 24GB | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ |
| 40GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
| 80GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
</details>
#### b. Weights
We'd better place the [weights](#model-zoo) along the specified path:
EasyAnimateV5:
```
📦 models/
├── 📂 Diffusion_Transformer/
│ └── 📂 PixArt-XL-2-512x512/
├── 📂 Motion_Module/
│ └── 📄 easyanimate_v1_mm.safetensors
├── 📂 Motion_Module/
│ ├── 📄 easyanimate_portrait.safetensors
│ └── 📄 easyanimate_portrait_lora.safetensors
│ ├── 📂 EasyAnimateV5-12b-zh-InP/
│ └── 📂 EasyAnimateV5-12b-zh/
├── 📂 Personalized_Model/
│ └── your trained trainformer model / your trained lora model (for UI load)
```
# 视频作品
The results displayed are all based on image.
### EasyAnimateV5-12b-zh-InP
#### I2V
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/bb393b7c-ba33-494c-ab06-b314adea9fc1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/cb0d0253-919d-4dd6-9dc1-5cd94443c7f1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/09ed361f-c0c5-4025-aad7-71fe1a1a52b1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/9f42848d-34eb-473f-97ea-a5ebd0268106" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/903fda91-a0bd-48ee-bf64-fff4e4d96f17" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/407c6628-9688-44b6-b12d-77de10fbbe95" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/ccf30ec1-91d2-4d82-9ce0-fcc585fc2f21" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/5dfe0f92-7d0d-43e0-b7df-0ff7b325663c" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/2b542b85-be19-4537-9607-9d28ea7e932e" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/c1662745-752d-4ad2-92bc-fe53734347b2" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/8bec3d66-50a3-4af5-a381-be2c865825a0" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/bcec22f4-732c-446f-958c-2ebbfd8f94be" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
#### T2V
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/eccb0797-4feb-48e9-91d3-5769ce30142b" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/76b3db64-9c7a-4d38-8854-dba940240ceb" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/0b8fab66-8de7-44ff-bd43-8f701bad6bb7" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/9fbddf5f-7fcd-4cc6-9d7c-3bdf1d4ce59e" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/19c1742b-e417-45ac-97d6-8bf3a80d8e13" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/641e56c8-a3d9-489d-a3a6-42c50a9aeca1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/2b16be76-518b-44c6-a69b-5c49d76df365" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/e7d9c0fc-136f-405c-9fab-629389e196be" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
### EasyAnimateV5-12b-zh-Control
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/53002ce2-dd18-4d4f-8135-b6f68364cabd" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/fce43c0b-81fa-4ab2-9ca7-78d786f520e6" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/b208b92c-5add-4ece-a200-3dbbe47b93c3" width="100%" controls autoplay loop></video>
</td>
<tr>
<td>
<video src="https://github.com/user-attachments/assets/3aec95d5-d240-49fb-a9e9-914446c7a4cf" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/60fa063b-5c1f-485f-b663-09bd6669de3f" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/4adde728-8397-42f3-8a2a-23f7b39e9a1e" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
# How to use
### 1. Inference
<h3 id="video-gen">1. Inference </h3>
#### a. Using Python Code
- Step 1: Download the corresponding weights and place them in the models folder.
- Step 1: Download the corresponding [weights](#model-zoo) and place them in the models folder.
- Step 2: Modify prompt, neg_prompt, guidance_scale, and seed in the predict_t2v.py file.
- Step 3: Run the predict_t2v.py file, wait for the generated results, and save the results in the samples/easyanimate-videos folder.
- Step 4: If you want to combine other backbones you have trained with Lora, modify the predict_t2v.py and Lora_path in predict_t2v.py depending on the situation.
#### b. Using webui
- Step 1: Download the corresponding weights and place them in the models folder.
- Step 2: Run the app. py file to enter the graph page.
- Step 1: Download the corresponding [weights](#model-zoo) and place them in the models folder.
- Step 2: Run the app.py file to enter the graph page.
- Step 3: Select the generated model based on the page, fill in prompt, neg_prompt, guidance_scale, and seed, click on generate, wait for the generated result, and save the result in the samples folder.
#### c. From ComfyUI
Please refer to [ComfyUI README](comfyui/README.md) for details.
#### d. GPU Memory Saving Schemes
Due to the large parameters of EasyAnimateV5, we need to consider GPU memory saving schemes to conserve memory. We provide a `GPU_memory_mode` option for each prediction file, which can be selected from `model_cpu_offload`, `model_cpu_offload_and_qfloat8`, and `sequential_cpu_offload`.
- `model_cpu_offload` indicates that the entire model will be offloaded to the CPU after use, saving some GPU memory.
- `model_cpu_offload_and_qfloat8` indicates that the entire model will be offloaded to the CPU after use, and the transformer model is quantized to float8, saving even more GPU memory.
- `sequential_cpu_offload` means that each layer of the model will be offloaded to the CPU after use, which is slower but saves a substantial amount of GPU memory.
### 2. Model Training
#### a、Training video generation model
##### i、Base on webvid dataset
If using the webvid dataset for training, you need to download the webvid dataset firstly.
A complete EasyAnimate training pipeline should include data preprocessing, Video VAE training, and Video DiT training. Among these, Video VAE training is optional because we have already provided a pre-trained Video VAE.
You need to arrange the webvid dataset in this format.
<h4 id="data-preprocess">a. data preprocessing</h4>
```
📦 project/
├── 📂 datasets/
│ ├── 📂 webvid/
│ ├── 📂 videos/
│ │ ├── 📄 00000001.mp4
│ │ ├── 📄 00000002.mp4
│ │ └── 📄 .....
│ └── 📄 csv_of_webvid.csv
```
We provide two simple demos:
- Train a Lora model using image data. For more details, you can refer to the [wiki](https://github.com/aigc-apps/EasyAnimate/wiki/Training-Lora).
- Perform SFT model training using video data. For more details, you can refer to the [wiki](https://github.com/aigc-apps/EasyAnimate/wiki/Training-SFT).
Then,set scripts/train_t2v.sh.
```
export DATASET_NAME="datasets/webvid/videos/"
export DATASET_META_NAME="datasets/webvid/csv_of_webvid.csv"
A complete data preprocessing link for long video segmentation, cleaning, and description can refer to [README](./easyanimate/video_caption/README.md) in the video captions section.
...
train_data_format="webvid"
```
Then, we run scripts/train_t2v.sh.
```sh
sh scripts/train_t2v.sh
```
##### ii、Base on internal dataset
If using the internal dataset for training, you need to format the dataset firstly.
You need to arrange the dataset in this format.
If you want to train a text to image and video generation model. You need to arrange the dataset in this format.
```
📦 project/
├── 📂 datasets/
│ ├── 📂 internal_datasets/
│ ├── 📂 videos/
│ ├── 📂 train/
│ │ ├── 📄 00000001.mp4
│ │ ├── 📄 00000002.mp4
│ │ ├── 📄 00000002.jpg
│ │ └── 📄 .....
│ └── 📄 json_of_internal_datasets.json
```
The json_of_internal_datasets.json is a standard JSON file, as shown in below:
The json_of_internal_datasets.json is a standard JSON file. The file_path in the json can to be set as relative path, as shown in below:
```json
[
{
"file_path": "videos/00000001.mp4",
"file_path": "train/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
{
"file_path": "videos/00000002.mp4",
"text": "A notepad with a drawing of a woman on it.",
"file_path": "train/00000002.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
.....
]
```
You can also set the path as absolute path as follow:
```json
[
{
"file_path": "/mnt/data/videos/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
}
.....
]
```
The file_path in the json needs to be set as relative path.
Then, set scripts/train_t2v.sh.
```
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/json_of_internal_datasets.json"
...
train_data_format="normal"
```
Then, we run scripts/train_t2v.sh.
```sh
sh scripts/train_t2v.sh
```
##### iii、Dataset Captioning (Optional)
We provide the captioning pipeline to get detailed descriptions of the video dataset. Please refer to [video_caption](./easyanimate/video_caption/) for details.
#### b、Training text to image model
##### i、Base on diffusers format
The format of dataset can be set as diffuser format.
If using the diffusers format dataset for training.
```
📦 project/
├── 📂 datasets/
│ ├── 📂 diffusers_datasets/
│ ├── 📂 train/
│ │ ├── 📄 00000001.jpg
│ │ ├── 📄 00000002.jpg
│ │ └── 📄 .....
│ └── 📄 metadata.jsonl
```
Then, set scripts/train_t2i.sh.
```
export DATASET_NAME="datasets/diffusers_datasets/"
...
train_data_format="diffusers"
```
Then, we run scripts/train_t2i.sh.
```sh
sh scripts/train_t2i.sh
```
##### ii、Base on internal dataset
If using the internal dataset for training, you need to format the dataset firstly.
You need to arrange the dataset in this format.
```
📦 project/
├── 📂 datasets/
│ ├── 📂 internal_datasets/
│ ├── 📂 train/
│ │ ├── 📄 00000001.jpg
│ │ ├── 📄 00000002.jpg
│ │ └── 📄 .....
│ └── 📄 json_of_internal_datasets.json
```
The json_of_internal_datasets.json is a standard JSON file, as shown in below:
```json
[
},
{
"file_path": "train/00000001.jpg",
"file_path": "/mnt/data/train/00000001.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
{
"file_path": "train/00000002.jpg",
"text": "A notepad with a drawing of a woman on it.",
"type": "image"
}
.....
]
```
The file_path in the json needs to be set as relative path.
Then, set scripts/train_t2i.sh.
<h4 id="vae-train">b. Video VAE training (optional)</h4>
Video VAE training is an optional option as we have already provided pre trained Video VAEs.
If you want to train video vae, you can refer to [README](easyanimate/vae/README.md) in the video vae section.
<h4 id="dit-train">c. Video DiT training </h4>
If the data format is relative path during data preprocessing, please set ```scripts/train.sh``` as follow.
```
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/json_of_internal_datasets.json"
...
train_data_format="normal"
```
Then, we run scripts/train_t2i.sh.
If the data format is absolute path during data preprocessing, please set ```scripts/train.sh``` as follow.
```
export DATASET_NAME=""
export DATASET_META_NAME="/mnt/data/json_of_internal_datasets.json"
```
Then, we run scripts/train.sh.
```sh
sh scripts/train_t2i.sh
sh scripts/train.sh
```
#### c、Training text to image Lora model
##### i、Base on diffusers format
The format of dataset can be set as diffuser format.
If using the diffusers format dataset for training.
For details on setting some parameters, please refer to [Readme Train](scripts/README_TRAIN.md) and [Readme Lora](scripts/README_TRAIN_LORA.md).
```
📦 project/
├── 📂 datasets/
│ ├── 📂 diffusers_datasets/
│ ├── 📂 train/
│ │ ├── 📄 00000001.jpg
│ │ ├── 📄 00000002.jpg
│ │ └── 📄 .....
│ └── 📄 metadata.jsonl
```
<details>
<summary>(Obsolete) EasyAnimateV1:</summary>
If you want to train EasyAnimateV1. Please switch to the git branch v1.
</details>
Then, set scripts/train_lora.sh.
```
export DATASET_NAME="datasets/diffusers_datasets/"
...
# Model zoo
train_data_format="diffusers"
```
EasyAnimateV5:
Then, we run scripts/train_lora.sh.
```sh
sh scripts/train_lora.sh
```
7B:
| Name | Type | Storage Space | Hugging Face | Model Scope | Description |
|--|--|--|--|--|--|
| EasyAnimateV5-7b-zh-InP | EasyAnimateV5 | 22 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-7b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-7b-zh-InP) | Official 7B image-to-video weights. Supports video prediction at multiple resolutions (512, 768, 1024), trained with 49 frames at 8 frames per second, and supports bilingual prediction in Chinese and English. |
| EasyAnimateV5-7b-zh | EasyAnimateV5 | 22 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-7b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-7b-zh) | Official 7B text-to-video weights. Supports video prediction at multiple resolutions (512, 768, 1024), trained with 49 frames at 8 frames per second, and supports bilingual prediction in Chinese and English. |
| EasyAnimateV5-Reward-LoRAs | EasyAnimateV5 | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-Reward-LoRAs) | The official reward backpropagation technology model optimizes the videos generated by EasyAnimateV5-12b to better match human preferences. |
##### ii、Base on internal dataset
If using the internal dataset for training, you need to format the dataset firstly.
12B:
| Name | Type | Storage Space | Hugging Face | Model Scope | Description |
|--|--|--|--|--|--|
| EasyAnimateV5-12b-zh-InP | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-InP) | Official image-to-video weights. Supports video prediction at multiple resolutions (512, 768, 1024), trained with 49 frames at 8 frames per second, and supports bilingual prediction in Chinese and English. |
| EasyAnimateV5-12b-zh-Control | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-Control) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-Control) | Official video control weights, supporting various control conditions such as Canny, Depth, Pose, MLSD, etc. Supports video prediction at multiple resolutions (512, 768, 1024) and is trained with 49 frames at 8 frames per second. Bilingual prediction in Chinese and English is supported. |
| EasyAnimateV5-12b-zh | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh) | Official text-to-video weights. Supports video prediction at multiple resolutions (512, 768, 1024), trained with 49 frames at 8 frames per second, and supports bilingual prediction in Chinese and English. |
| EasyAnimateV5-Reward-LoRAs | EasyAnimateV5 | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-Reward-LoRAs) | The official reward backpropagation technology model optimizes the videos generated by EasyAnimateV5-12b to better match human preferences. |
You need to arrange the dataset in this format.
<details>
<summary>(Obsolete) EasyAnimateV4:</summary>
```
📦 project/
├── 📂 datasets/
│ ├── 📂 internal_datasets/
│ ├── 📂 train/
│ │ ├── 📄 00000001.jpg
│ │ ├── 📄 00000002.jpg
│ │ └── 📄 .....
│ └── 📄 json_of_internal_datasets.json
```
| Name | Type | Storage Space | Hugging Face | Model Scope | Description |
|--|--|--|--|--|--|
| EasyAnimateV4-XL-2-InP.tar.gz | EasyAnimateV4 | Before extraction: 8.9 GB \/ After extraction: 14.0 GB |[🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV4-XL-2-InP)| [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV4-XL-2-InP)| | Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 144 frames at a rate of 24 frames per second. |
</details>
The json_of_internal_datasets.json is a standard JSON file, as shown in below:
```json
[
{
"file_path": "train/00000001.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
{
"file_path": "train/00000002.jpg",
"text": "A notepad with a drawing of a woman on it.",
"type": "image"
}
.....
]
```
The file_path in the json needs to be set as relative path.
<details>
<summary>(Obsolete) EasyAnimateV3:</summary>
Then, set scripts/train_lora.sh.
```
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/json_of_internal_datasets.json"
| Name | Type | Storage Space | Hugging Face | Model Scope | Description |
|--|--|--|--|--|--|
| EasyAnimateV3-XL-2-InP-512x512.tar | EasyAnimateV3 | 18.2GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-512x512)| [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV3-XL-2-InP-512x512) | EasyAnimateV3 official weights for 512x512 text and image to video resolution. Training with 144 frames and fps 24 |
| EasyAnimateV3-XL-2-InP-768x768.tar | EasyAnimateV3 | 18.2GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-768x768) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV3-XL-2-InP-768x768) | EasyAnimateV3 official weights for 768x768 text and image to video resolution. Training with 144 frames and fps 24 |
| EasyAnimateV3-XL-2-InP-960x960.tar | EasyAnimateV3 | 18.2GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-960x960) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV3-XL-2-InP-960x960) | EasyAnimateV3 official weights for 960x960 text and image to video resolution. Training with 144 frames and fps 24 |
</details>
...
<details>
<summary>(Obsolete) EasyAnimateV2:</summary>
train_data_format="normal"
```
| Name | Type | Storage Space | Url | Hugging Face | Model Scope | Description |
|--|--|--|--|--|--|--|
| EasyAnimateV2-XL-2-512x512.tar | EasyAnimateV2 | 16.2GB | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV2-XL-2-512x512)| [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV2-XL-2-512x512)| EasyAnimateV2 official weights for 512x512 resolution. Training with 144 frames and fps 24 |
| EasyAnimateV2-XL-2-768x768.tar | EasyAnimateV2 | 16.2GB | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV2-XL-2-768x768) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV2-XL-2-768x768)| EasyAnimateV2 official weights for 768x768 resolution. Training with 144 frames and fps 24 |
| easyanimatev2_minimalism_lora.safetensors | Lora of Pixart | 485.1MB | [Download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimatev2_minimalism_lora.safetensors)| - | - | A lora training with a specifial type images. Images can be downloaded from [Url](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v2/Minimalism.zip). |
</details>
Then, we run scripts/train_lora.sh.
```sh
sh scripts/train_lora.sh
```
<details>
<summary>(Obsolete) EasyAnimateV1:</summary>
# Algorithm Detailed
We build EasyAnimate by introducing additional motion module upon [PixArt-alpha](https://github.com/PixArt-alpha/PixArt-alpha),so that can extend the DiT model from 2D image generation to 3D video generation. The pipeline is shwon as follows.
### 1、Motion Weights
| Name | Type | Storage Space | Url | Description |
|--|--|--|--|--|
| easyanimate_v1_mm.safetensors | Motion Module | 4.1GB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Motion_Module/easyanimate_v1_mm.safetensors) | Training with 80 frames and fps 12 |
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/pipeline.png" alt="ui" style="zoom:50%;" />
### 2、Other Weights
| Name | Type | Storage Space | Url | Description |
|--|--|--|--|--|
| PixArt-XL-2-512x512.tar | Pixart | 11.4GB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Diffusion_Transformer/PixArt-XL-2-512x512.tar)| Pixart-Alpha official weights |
| easyanimate_portrait.safetensors | Checkpoint of Pixart | 2.3GB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait.safetensors) | Training with internal portrait datasets |
| easyanimate_portrait_lora.safetensors | Lora of Pixart | 654.0MB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait_lora.safetensors)| Training with internal portrait datasets |
</details>
The motion module is used to capture the temporal information among frames. The structure is shown as follows.
# TODO List
- Support model with larger params.
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/motion_module.png" alt="motion" style="zoom:50%;" />
# Contact Us
1. Use Dingding to search group 77450006752 or Scan to join
2. You need to scan the image to join the WeChat group or if it is expired, add this student as a friend first to invite you.
We introduce attention mechanisms in the temporal dimension to enable the model to learn temporal information for generating continuous video frames. At the same time, we utilize an additional Grid Reshape calculation to expand the number of input tokens for the attention mechanism, thus making greater use of the spatial information in images to achieve better generative results.
The Motion Module, as a separate module, can be applied to different DiT baseline models during inference. Furthermore, EasyAnimate not only supports the training of the motion-module but also supports the training of the DiT base model/LoRA model, making it convenient for users to complete training of a customized-style model according to their own needs and thereby generate videos of any style.
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/group/dd.png" alt="ding group" width="30%"/>
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/group/wechat.jpg" alt="Wechat group" width="30%"/>
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/group/person.jpg" alt="Person" width="30%"/>
# Reference
- CogVideo: https://github.com/THUDM/CogVideo/
- Flux: https://github.com/black-forest-labs/flux
- magvit: https://github.com/google-research/magvit
- PixArt: https://github.com/PixArt-alpha/PixArt-alpha
- Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
- Open-Sora: https://github.com/hpcaitech/Open-Sora
- Animatediff: https://github.com/guoyww/AnimateDiff
- ComfyUI-EasyAnimateWrapper: https://github.com/kijai/ComfyUI-EasyAnimateWrapper
- HunYuan DiT: https://github.com/tencent/HunyuanDiT
# License
This project is licensed under the [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE).
This project is licensed under the [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE).
+492
View File
@@ -0,0 +1,492 @@
# 📷 EasyAnimate | 高解像度および長時間動画生成のためのエンドツーエンドソリューション
😊 EasyAnimateは、高解像度および長時間動画を生成するためのエンドツーエンドソリューションです。トランスフォーマーベースの拡散生成器をトレーニングし、長時間動画を処理するためのVAEをトレーニングし、メタデータを前処理することができます。
😊 DITをベースに、トランスフォーマーを拡散器として使用して動画や画像を生成します。
😊 ようこそ!
[![Arxiv Page](https://img.shields.io/badge/Arxiv-Page-red)](https://arxiv.org/abs/2405.18991)
[![Project Page](https://img.shields.io/badge/Project-Website-green)](https://easyanimate.github.io/)
[![Modelscope Studio](https://img.shields.io/badge/Modelscope-Studio-blue)](https://modelscope.cn/studios/PAI/EasyAnimate/summary)
[![Hugging Face Spaces](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Spaces-yellow)](https://huggingface.co/spaces/alibaba-pai/EasyAnimate)
[![Discord Page](https://img.shields.io/badge/Discord-Page-blue)](https://discord.gg/UzkpB4Bn)
English | [简体中文](./README_zh-CN.md) | 日本語
# 目次
- [目次](#目次)
- [紹介](#紹介)
- [クイックスタート](#クイックスタート)
- [ビデオ結果](#ビデオ結果)
- [使い方](#使い方)
- [モデルズー](#モデルズー)
- [TODOリスト](#todoリスト)
- [お問い合わせ](#お問い合わせ)
- [参考文献](#参考文献)
- [ライセンス](#ライセンス)
# 紹介
EasyAnimateは、トランスフォーマーアーキテクチャに基づいたパイプラインで、AI画像および動画の生成、Diffusion TransformerのベースラインモデルおよびLoraモデルのトレーニングに使用されます。事前トレーニング済みのEasyAnimateモデルから直接予測を行い、さまざまな解像度で約6秒間、8fpsの動画を生成できます(EasyAnimateV5、1〜49フレーム)。さらに、ユーザーは特定のスタイル変換のために独自のベースラインおよびLoraモデルをトレーニングできます。
異なるプラットフォームからのクイックプルアップをサポートします。詳細は[クイックスタート](#クイックスタート)を参照してください。
**新機能:**
- インセンティブ逆伝播を使用してLoraを訓練し、人間の好みに合うようにビデオを最適化します。詳細は、[ここ](scripts/README _ train _ REVARD.md)を参照してください。EasyAnimateV 5-7 bがリリースされました。[2024.11.27]
- **v5に更新**、1024x1024までの動画生成をサポート、49フレーム、6秒、8fps、モデルスケールを12Bに拡張、MMDIT構造を組み込み、さまざまな入力を持つ制御モデルをサポート。中国語と英語のバイリンガル予測をサポート。[2024.11.08]
- **v4に更新**、1024x1024までの動画生成をサポート、144フレーム、6秒、24fps、テキスト、画像、動画からの動画生成をサポート、512から1280までの解像度を単一モデルで処理。中国語と英語のバイリンガル予測をサポート。[2024.08.15]
- **v3に更新**、960x960までの動画生成をサポート、144フレーム、6秒、24fps、テキストと画像からの動画生成をサポート。[2024.07.01]
- **ModelScope-Sora “データディレクター” クリエイティブレース** — 第三回Data-Juicerビッグモデルデータチャレンジが正式に開始されました!EasyAnimateをベースモデルとして使用し、データ処理がモデルトレーニングに与える影響を探ります。詳細は[競技ウェブサイト](https://tianchi.aliyun.com/competition/entrance/532219)をご覧ください。[2024.06.17]
- **v2に更新**、768x768までの動画生成をサポート、144フレーム、6秒、24fps。[2024.05.26]
- **コード作成!** 現在、WindowsおよびLinuxをサポート。[2024.04.12]
機能:
- [データ前処理](#data-preprocess)
- [VAEのトレーニング](#vae-train)
- [DiTのトレーニング](#dit-train)
- [動画生成](#video-gen)
私たちのUIインターフェースは次のとおりです:
![ui](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/ui_v3.jpg)
# クイックスタート
### 1. クラウド使用: AliyunDSW/Docker
#### a. AliyunDSWから
DSWには無料のGPU時間があり、ユーザーは一度申請でき、申請後3ヶ月間有効です。
Aliyunは[Freetier](https://free.aliyun.com/?product=9602825&crowd=enterprise&spm=5176.28055625.J_5831864660.1.e939154aRgha4e&scm=20140722.M_9974135.P_110.MO_1806-ID_9974135-MID_9974135-CID_30683-ST_8512-V_1)で無料のGPU時間を提供しており、取得してAliyun PAI-DSWで使用し、5分以内にEasyAnimateを開始できます!
[![DSW Notebook](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/dsw.png)](https://gallery.pai-ml.com/#/preview/deepLearning/cv/easyanimate_v5)
#### b. ComfyUIから
私たちのComfyUIは次のとおりです。詳細は[ComfyUI README](comfyui/README.md)を参照してください。
![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v3/comfyui_i2v.jpg)
#### c. Dockerから
Dockerを使用している場合は、マシンにグラフィックスカードドライバとCUDA環境が正しくインストールされていることを確認してください。
次のコマンドを実行します:
```
# イメージをプル
docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:easyanimate
# イメージに入る
docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:easyanimate
# コードをクローン
git clone https://github.com/aigc-apps/EasyAnimate.git
# EasyAnimateのディレクトリに入る
cd EasyAnimate
# 重みをダウンロード
mkdir models/Diffusion_Transformer
mkdir models/Motion_Module
mkdir models/Personalized_Model
# EasyAnimateV5モデルをダウンロードするには、hugginfaceリンクまたはmodelscopeリンクを使用してください。
# I2Vモデル
# https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-InP
# https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-InP
# T2Vモデル
# https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh
# https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh
```
### 2. ローカルインストール: 環境チェック/ダウンロード/インストール
#### a. 環境チェック
次の環境でEasyAnimateの実行を確認しました:
Windowsの詳細:
- OS: Windows 10
- python: python3.10 & python3.11
- pytorch: torch2.2.0
- CUDA: 11.8 & 12.1
- CUDNN: 8+
- GPU: Nvidia-3060 12G
Linuxの詳細:
- OS: Ubuntu 20.04, CentOS
- python: python3.10 & python3.11
- pytorch: torch2.2.0
- CUDA: 11.8 & 12.1
- CUDNN: 8+
- GPU:Nvidia-V100 16G & Nvidia-A10 24G & Nvidia-A100 40G & Nvidia-A100 80G
ディスクに約60GBの空き容量が必要です(重みを保存するため)、確認してください!
EasyAnimateV5-12Bのビデオサイズは異なるGPUメモリにより生成できます。以下の表をご覧ください:
| GPU memory |384x672x25|384x672x49|576x1008x25|576x1008x49|768x1344x25|768x1344x49|
|----------|----------|----------|----------|----------|----------|----------|
| 16GB | 🧡 | 🧡 | ❌ | ❌ | ❌ | ❌ |
| 24GB | 🧡 | 🧡 | 🧡 | 🧡 | 🧡 | ❌ |
| 40GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
| 80GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
EasyAnimateV5-7Bのビデオサイズは異なるGPUメモリにより生成できます。以下の表をご覧ください:
| GPU memory |384x672x25|384x672x49|576x1008x25|576x1008x49|768x1344x25|768x1344x49|
|----------|----------|----------|----------|----------|----------|----------|
| 16GB | 🧡 | 🧡 | ❌ | ❌ | ❌ | ❌ |
| 24GB | ✅ | ✅ | ✅ | 🧡 | 🧡 | ❌ |
| 40GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
| 80GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
✅ は"model_cpu_offload"の条件で実行可能であることを示し、🧡は"model_cpu_offload_and_qfloat8"の条件で実行可能を示し、⭕️ は"sequential_cpu_offload"の条件では実行可能であることを示しています。❌は実行できないことを示します。sequential_cpu_offloadにより実行する場合は遅くなります。
一部のGPU(例:2080ti、V100)はtorch.bfloat16をサポートしていないため、app.pyおよびpredictファイル内のweight_dtypeをtorch.float16に変更する必要があります。
EasyAnimateV5-12Bは異なるGPUで25ステップ生成する時間は次の通りです:
| GPU |384x672x25|384x672x49|576x1008x25|576x1008x49|768x1344x25|768x1344x49|
|----------|----------|----------|----------|----------|----------|----------|
| A10 24GB |約120秒 (4.8s/it)|約240秒 (9.6s/it)|約320秒 (12.7s/it)|約750秒 (29.8s/it)| ❌ | ❌ |
| A100 80GB |約45秒 (1.75s/it)|約90秒 (3.7s/it)|約120秒 (4.7s/it)|約300秒 (11.4s/it)|約265秒 (10.6s/it)| 約710秒 (28.3s/it)|
(⭕️) はlow_gpu_memory_mode=Trueの条件で実行可能であるが、速度が遅くなることを示しています。また、❌は実行できないことを示します。
<details>
<summary>(廃止予定) EasyAnimateV3:</summary>
EasyAnimateV3のビデオサイズは異なるGPUメモリにより生成できます。以下の表をご覧ください:
| GPUメモリ | 384x672x25 | 384x672x144 | 576x1008x72 | 576x1008x144 | 720x1280x72 | 720x1280x144 |
|----------|----------|----------|----------|----------|----------|----------|
| 12GB | ⭕️ | ⭕️ | ⭕️ | ⭕️ | ❌ | ❌ |
| 16GB | ✅ | ✅ | ⭕️ | ⭕️ | ⭕️ | ❌ |
| 24GB | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ |
| 40GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
| 80GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
</details>
#### b. 重み
[重み](#model-zoo)を指定されたパスに配置することをお勧めします:
EasyAnimateV5:
```
📦 models/
├── 📂 Diffusion_Transformer/
│ ├── 📂 EasyAnimateV5-12b-zh-InP/
│ └── 📂 EasyAnimateV5-12b-zh/
├── 📂 Personalized_Model/
│ └── あなたのトレーニング済みのトランスフォーマーモデル / あなたのトレーニング済みのLoraモデル(UIロード用)
```
# ビデオ結果
表示されている結果はすべて画像からの生成に基づいています。
### EasyAnimateV5-12b-zh-InP
#### I2V
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/bb393b7c-ba33-494c-ab06-b314adea9fc1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/cb0d0253-919d-4dd6-9dc1-5cd94443c7f1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/09ed361f-c0c5-4025-aad7-71fe1a1a52b1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/9f42848d-34eb-473f-97ea-a5ebd0268106" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/903fda91-a0bd-48ee-bf64-fff4e4d96f17" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/407c6628-9688-44b6-b12d-77de10fbbe95" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/ccf30ec1-91d2-4d82-9ce0-fcc585fc2f21" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/5dfe0f92-7d0d-43e0-b7df-0ff7b325663c" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/2b542b85-be19-4537-9607-9d28ea7e932e" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/c1662745-752d-4ad2-92bc-fe53734347b2" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/8bec3d66-50a3-4af5-a381-be2c865825a0" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/bcec22f4-732c-446f-958c-2ebbfd8f94be" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
#### T2V
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/eccb0797-4feb-48e9-91d3-5769ce30142b" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/76b3db64-9c7a-4d38-8854-dba940240ceb" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/0b8fab66-8de7-44ff-bd43-8f701bad6bb7" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/9fbddf5f-7fcd-4cc6-9d7c-3bdf1d4ce59e" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/19c1742b-e417-45ac-97d6-8bf3a80d8e13" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/641e56c8-a3d9-489d-a3a6-42c50a9aeca1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/2b16be76-518b-44c6-a69b-5c49d76df365" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/e7d9c0fc-136f-405c-9fab-629389e196be" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
### EasyAnimateV5-12b-zh-Control
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/53002ce2-dd18-4d4f-8135-b6f68364cabd" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/fce43c0b-81fa-4ab2-9ca7-78d786f520e6" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/b208b92c-5add-4ece-a200-3dbbe47b93c3" width="100%" controls autoplay loop></video>
</td>
<tr>
<td>
<video src="https://github.com/user-attachments/assets/3aec95d5-d240-49fb-a9e9-914446c7a4cf" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/60fa063b-5c1f-485f-b663-09bd6669de3f" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/4adde728-8397-42f3-8a2a-23f7b39e9a1e" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
# 使い方
<h3 id="video-gen">1. 推論 </h3>
#### a. Pythonコードを使用する
- ステップ1:対応する[重み](#model-zoo)をダウンロードし、modelsフォルダに配置します。
- ステップ2:predict_t2v.pyファイルでprompt、neg_prompt、guidance_scale、およびseedを変更します。
- ステップ3:predict_t2v.pyファイルを実行し、生成された結果を待ちます。結果はsamples/easyanimate-videosフォルダに保存されます。
- ステップ4:他のバックボーンとLoraを組み合わせたい場合は、状況に応じてpredict_t2v.pyおよびLora_pathを変更します。
#### b. WebUIを使用する
- ステップ1:対応する[重み](#model-zoo)をダウンロードし、modelsフォルダに配置します。
- ステップ2:app.pyファイルを実行してグラフページに入ります。
- ステップ3:ページに基づいて生成モデルを選択し、prompt、neg_prompt、guidance_scale、およびseedを入力し、生成をクリックして生成結果を待ちます。結果はsamplesフォルダに保存されます。
#### c. ComfyUIから
詳細は[ComfyUI README](comfyui/README.md)を参照してください。
#### d. GPUメモリ節約スキーム
EasyAnimateV5のパラメータが大きいため、メモリを節約するためにGPUメモリ節約スキームを検討する必要があります。各予測ファイルには、`GPU_memory_mode`オプションがあり、`model_cpu_offload`、`model_cpu_offload_and_qfloat8`、および`sequential_cpu_offload`から選択できます。
- `model_cpu_offload`は、使用後にモデル全体がCPUにオフロードされることを示し、一部のGPUメモリを節約します。
- `model_cpu_offload_and_qfloat8`は、使用後にモデル全体がCPUにオフロードされ、トランスフォーマーモデルがfloat8に量子化され、さらに多くのGPUメモリを節約します。
- `sequential_cpu_offload`は、使用後にモデルの各層がCPUにオフロードされることを意味し、速度は遅くなりますが、大量のGPUメモリを節約します。
### 2. モデルトレーニング
完全なEasyAnimateトレーニングパイプラインには、データ前処理、Video VAEトレーニング、およびVideo DiTトレーニングが含まれる必要があります。これらの中で、Video VAEトレーニングはオプションです。すでにトレーニング済みのVideo VAEを提供しているためです。
<h4 id="data-preprocess">a. データ前処理</h4>
私たちは2つの簡単なデモを提供します:
- 画像データを使用してLoraモデルを訓練します。詳細はこちらの[wiki](https://github.com/aigc-apps/EasyAnimate/wiki/Training-Lora)をご覧ください。
- 動画データを使用してSFTモデルを訓練します。詳細はこちらの[wiki](https://github.com/aigc-apps/EasyAnimate/wiki/Training-SFT)をご覧ください。
長時間動画のセグメンテーション、クリーニング、および説明のための完全なデータ前処理リンクは、ビデオキャプションセクションの[README](./easyanimate/video_caption/README.md)を参照してください。
テキストから画像および動画生成モデルをトレーニングする場合は、データセットを次の形式で配置する必要があります。
```
📦 project/
├── 📂 datasets/
│ ├── 📂 internal_datasets/
│ ├── 📂 train/
│ │ ├── 📄 00000001.mp4
│ │ ├── 📄 00000002.jpg
│ │ └── 📄 .....
│ └── 📄 json_of_internal_datasets.json
```
json_of_internal_datasets.jsonは標準のJSONファイルです。json内のfile_pathは相対パスとして設定できます。以下のように:
```json
[
{
"file_path": "train/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
{
"file_path": "train/00000002.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
.....
]
```
パスを絶対パスとして設定することもできます:
```json
[
{
"file_path": "/mnt/data/videos/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
{
"file_path": "/mnt/data/train/00000001.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
.....
]
```
<h4 id="vae-train">b. Video VAEトレーニング(オプション)</h4>
Video VAEトレーニングはオプションです。すでにトレーニング済みのVideo VAEを提供しているためです。
Video VAEをトレーニングする場合は、ビデオVAEセクションの[README](easyanimate/vae/README.md)を参照してください。
<h4 id="dit-train">c. Video DiTトレーニング </h4>
データ前処理時にデータ形式が相対パスの場合、```scripts/train.sh```を次のように設定します。
```
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/json_of_internal_datasets.json"
```
データ前処理時にデータ形式が絶対パスの場合、```scripts/train.sh```を次のように設定します。
```
export DATASET_NAME=""
export DATASET_META_NAME="/mnt/data/json_of_internal_datasets.json"
```
次に、scripts/train.shを実行します。
```sh
sh scripts/train.sh
```
一部のパラメータの設定の詳細については、[Readme Train](scripts/README_TRAIN.md)および[Readme Lora](scripts/README_TRAIN_LORA.md)を参照してください。
<details>
<summary>(Obsolete) EasyAnimateV1:</summary>
EasyAnimateV1をトレーニングする場合は、gitブランチv1に切り替えてください。
</details>
# モデルズー
EasyAnimateV5:
7B:
| 名前 | 種類 | ストレージスペース | Hugging Face | Model Scope | 説明 |
|--|--|--|--|--|--|
| EasyAnimateV5-7b-zh-InP | EasyAnimateV5 | 22 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-7b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-7b-zh-InP) | 公式の画像から動画への重み。複数の解像度(512、768、1024)での動画予測をサポートし、49フレーム、毎秒8フレームでトレーニングされ、中国語と英語のバイリンガル予測をサポートします。 |
| EasyAnimateV5-7b-zh | EasyAnimateV5 | 22 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-7b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-7b-zh) | 公式のテキストから動画への重み。複数の解像度(512、768、1024)での動画予測をサポートし、49フレーム、毎秒8フレームでトレーニングされ、中国語と英語のバイリンガル予測をサポートします。 |
| EasyAnimateV5-Reward-LoRAs | EasyAnimateV5 | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-Reward-LoRAs) | 公式インバース伝播技術モデルによるEasyAnimateV 5-12 b生成ビデオの最適化によるヒト選好の最適化|
12B:
| 名前 | 種類 | ストレージスペース | Hugging Face | Model Scope | 説明 |
|--|--|--|--|--|--|
| EasyAnimateV5-12b-zh-InP | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-InP) | 公式の画像から動画への重み。複数の解像度(512、768、1024)での動画予測をサポートし、49フレーム、毎秒8フレームでトレーニングされ、中国語と英語のバイリンガル予測をサポートします。 |
| EasyAnimateV5-12b-zh-Control | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-Control) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-Control) | 公式の動画制御重み。Canny、Depth、Pose、MLSDなどのさまざまな制御条件をサポートします。複数の解像度(512、768、1024)での動画予測をサポートし、49フレーム、毎秒8フレームでトレーニングされ、中国語と英語のバイリンガル予測をサポートします。 |
| EasyAnimateV5-12b-zh | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh) | 公式のテキストから動画への重み。複数の解像度(512、768、1024)での動画予測をサポートし、49フレーム、毎秒8フレームでトレーニングされ、中国語と英語のバイリンガル予測をサポートします。 |
| EasyAnimateV5-Reward-LoRAs | EasyAnimateV5 | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-Reward-LoRAs) | 公式インバース伝播技術モデルによるEasyAnimateV 5-12 b生成ビデオの最適化によるヒト選好の最適化|
<details>
<summary>(Obsolete) EasyAnimateV4:</summary>
| 名前 | 種類 | ストレージスペース | Hugging Face | Model Scope | 説明 |
|--|--|--|--|--|--|
| EasyAnimateV4-XL-2-InP.tar.gz | EasyAnimateV4 | 解凍前: 8.9 GB / 解凍後: 14.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV4-XL-2-InP)| [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV4-XL-2-InP) | 公式のグラフ生成動画モデル。複数の解像度(512、768、1024、1280)での動画予測をサポートし、144フレーム、毎秒24フレームでトレーニングされています。 |
</details>
<details>
<summary>(Obsolete) EasyAnimateV3:</summary>
| 名前 | 種類 | ストレージスペース | Hugging Face | Model Scope | 説明 |
|--|--|--|--|--|--|
| EasyAnimateV3-XL-2-InP-512x512.tar | EasyAnimateV3 | 18.2GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-512x512)| [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV3-XL-2-InP-512x512) | EasyAnimateV3公式の512x512テキストおよび画像から動画への重み。144フレーム、毎秒24フレームでトレーニングされています。 |
| EasyAnimateV3-XL-2-InP-768x768.tar | EasyAnimateV3 | 18.2GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-768x768) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV3-XL-2-InP-768x768) | EasyAnimateV3公式の768x768テキストおよび画像から動画への重み。144フレーム、毎秒24フレームでトレーニングされています。 |
| EasyAnimateV3-XL-2-InP-960x960.tar | EasyAnimateV3 | 18.2GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-960x960) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV3-XL-2-InP-960x960) | EasyAnimateV3公式の960x960テキストおよび画像から動画への重み。144フレーム、毎秒24フレームでトレーニングされています。 |
</details>
<details>
<summary>(Obsolete) EasyAnimateV2:</summary>
| 名前 | 種類 | ストレージスペース | URL | Hugging Face | Model Scope | 説明 |
|--|--|--|--|--|--|--|
| EasyAnimateV2-XL-2-512x512.tar | EasyAnimateV2 | 16.2GB | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV2-XL-2-512x512)| [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV2-XL-2-512x512) | EasyAnimateV2公式の512x512解像度の重み。144フレーム、毎秒24フレームでトレーニングされています。 |
| EasyAnimateV2-XL-2-768x768.tar | EasyAnimateV2 | 16.2GB | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV2-XL-2-768x768) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV2-XL-2-768x768) | EasyAnimateV2公式の768x768解像度の重み。144フレーム、毎秒24フレームでトレーニングされています。 |
| easyanimatev2_minimalism_lora.safetensors | Lora of Pixart | 485.1MB | [ダウンロード](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimatev2_minimalism_lora.safetensors) | - | - | 特定のタイプの画像でトレーニングされたLora。画像は[URL](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v2/Minimalism.zip)からダウンロードできます。 |
</details>
<details>
<summary>(Obsolete) EasyAnimateV1:</summary>
### 1、モーション重み
| 名前 | 種類 | ストレージスペース | URL | 説明 |
|--|--|--|--|--|
| easyanimate_v1_mm.safetensors | モーションモジュール | 4.1GB | [ダウンロード](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Motion_Module/easyanimate_v1_mm.safetensors) | 80フレーム、毎秒12フレームでトレーニングされています。 |
### 2、その他の重み
| 名前 | 種類 | ストレージスペース | URL | 説明 |
|--|--|--|--|--|
| PixArt-XL-2-512x512.tar | Pixart | 11.4GB | [ダウンロード](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Diffusion_Transformer/PixArt-XL-2-512x512.tar)| Pixart-Alpha公式の重み。 |
| easyanimate_portrait.safetensors | Pixartのチェックポイント | 2.3GB | [ダウンロード](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait.safetensors) | 内部のポートレートデータセットでトレーニングされています。 |
| easyanimate_portrait_lora.safetensors | PixartのLora | 654.0MB | [ダウンロード](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait_lora.safetensors)| 内部のポートレートデータセットでトレーニングされています。 |
</details>
# TODOリスト
- より大きなパラメータを持つモデルをサポートします。
# お問い合わせ
1. Dingdingを使用してグループ77450006752を検索するか、スキャンして参加します。
2. WeChatグループに参加するには画像をスキャンするか、期限切れの場合はこの学生を友達として追加して招待します。
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/group/dd.png" alt="ding group" width="30%"/>
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/group/wechat.jpg" alt="Wechat group" width="30%"/>
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/group/person.jpg" alt="Person" width="30%"/>
# 参考文献
- CogVideo: https://github.com/THUDM/CogVideo/
- Flux: https://github.com/black-forest-labs/flux
- magvit: https://github.com/google-research/magvit
- PixArt: https://github.com/PixArt-alpha/PixArt-alpha
- Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
- Open-Sora: https://github.com/hpcaitech/Open-Sora
- Animatediff: https://github.com/guoyww/AnimateDiff
- ComfyUI-EasyAnimateWrapper: https://github.com/kijai/ComfyUI-EasyAnimateWrapper
- HunYuan DiT: https://github.com/tencent/HunyuanDiT
# ライセンス
このプロジェクトは[Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE)の下でライセンスされています。
+444 -395
View File
@@ -1,52 +1,460 @@
# EasyAnimate | 您的智能生成器。
😊 EasyAnimate是一个用于生成长视频和训练基于transformer的扩散生成器的repo。
# EasyAnimate | 高分辨率长视频生成的端到端解决方案
😊 EasyAnimate是一个用于生成高分辨率和长视频的端到端解决方案。我们可以训练基于转换器的扩散生成器,训练用于处理长视频的VAE,以及预处理元数据。
😊 我们基于类SORA结构与DIT,使用transformer进行作为扩散器进行视频生成。为了保证良好的拓展性,我们基于motion module构建了EasyAnimate,未来我们也会尝试更多的训练方案一提高效果。
😊 我们基于DIT,使用transformer进行作为扩散器进行视频与图片生成。
😊 Welcome!
[![Arxiv Page](https://img.shields.io/badge/Arxiv-Page-red)](https://arxiv.org/abs/2405.18991)
[![Project Page](https://img.shields.io/badge/Project-Website-green)](https://easyanimate.github.io/)
[![Modelscope Studio](https://img.shields.io/badge/Modelscope-Studio-blue)](https://modelscope.cn/studios/PAI/EasyAnimate/summary)
[![Hugging Face Spaces](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Spaces-yellow)](https://huggingface.co/spaces/alibaba-pai/EasyAnimate)
[![Discord Page](https://img.shields.io/badge/Discord-Page-blue)](https://discord.gg/UzkpB4Bn)
[English](./README.md) | 简体中文
[English](./README.md) | 简体中文 | [日本語](./README_ja-JP.md)
# 目录
- [目录](#目录)
- [简介](#简介)
- [TODO List](#todo-list)
- [Model zoo](#model-zoo)
- [1、运动权重](#1运动权重)
- [2、其他权重](#2其他权重)
- [快速启动](#快速启动)
- [1. 云使用: AliyunDSW/Docker](#1-云使用-aliyundswdocker)
- [2. 本地安装: 环境检查/下载/安装](#2-本地安装-环境检查下载安装)
- [视频作品](#视频作品)
- [如何使用](#如何使用)
- [1. 生成](#1-生成)
- [2. 模型训练](#2-模型训练)
- [算法细节](#算法细节)
- [模型地址](#模型地址)
- [未来计划](#未来计划)
- [联系我们](#联系我们)
- [参考文献](#参考文献)
- [许可证](#许可证)
# 简介
EasyAnimate是一个基于transformer结构的pipeline,可用于生成AI动画、训练Diffusion Transformer的基线模型与Lora模型,我们支持从已经训练好的EasyAnimate模型直接进行预测,生成不同分辨率,6秒左右、fps12的视频(40 ~ 80帧, 未来会支持更长的视频),也支持用户训练自己的基线模型与Lora模型,进行一定的风格变换。
EasyAnimate是一个基于transformer结构的pipeline,可用于生成AI图片与视频、训练Diffusion Transformer的基线模型与Lora模型,我们支持从已经训练好的EasyAnimate模型直接进行预测,生成不同分辨率,6秒左右、fps8的视频(EasyAnimateV5,1 ~ 49帧),也支持用户训练自己的基线模型与Lora模型,进行一定的风格变换。
我们会逐渐支持从不同平台快速启动,请参阅 [快速启动](#快速启动)。
新特性:
- 添加视频数据标注的[代码](./easyanimate/video_caption/)。[ 2024.04.17 ]
- 使用奖励反向传播来训练Lora并优化视频,使其更好地符合人类偏好,详细信息请参见[此处](scripts/README_train_REVARD.md)。EasyAnimateV5-7b现已发布。[ 2024.11.27 ]
- 更新到v5版本,最大支持1024x1024,49帧, 6s, 8fps视频生成,拓展模型规模到12B,应用MMDIT结构,支持不同输入的控制模型,支持中文与英文双语预测。[ 2024.11.08 ]
- 更新到v4版本,最大支持1024x1024,144帧, 6s, 24fps视频生成,支持文、图、视频生视频,单个模型可支持512到1280任意分辨率,支持中文与英文双语预测。[ 2024.08.15 ]
- 更新到v3版本,最大支持960x960,144帧,6s, 24fps视频生成,支持文与图生视频模型。[ 2024.07.01 ]
- ModelScope-Sora“数据导演”创意竞速——第三届Data-Juicer大模型数据挑战赛已经正式启动!其使用EasyAnimate作为基础模型,探究数据处理对于模型训练的作用。立即访问[竞赛官网](https://tianchi.aliyun.com/competition/entrance/532219),了解赛事详情。[ 2024.06.17 ]
- 更新到v2版本,最大支持768x768,144帧,6s, 24fps视频生成。[ 2024.05.26 ]
- 创建代码!现在支持 Windows 和 Linux。[ 2024.04.12 ]
这些是我们的生成结果:
功能概览:
- [数据预处理](#data-preprocess)
- [训练VAE](#vae-train)
- [训练DiT](#dit-train)
- [模型生成](#video-gen)
我们的ui界面如下:
![ui](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/ui.png)
![ui](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/ui_v3.jpg)
# TODO List
- 支持更大分辨率的文视频生成模型。
- 支持基于magvit的文视频生成模型。
- 支持视频inpaint模型。
# 快速启动
### 1. 云使用: AliyunDSW/Docker
#### a. 通过阿里云 DSW
DSW 有免费 GPU 时间,用户可申请一次,申请后3个月内有效。
阿里云在[Freetier](https://free.aliyun.com/?product=9602825&crowd=enterprise&spm=5176.28055625.J_5831864660.1.e939154aRgha4e&scm=20140722.M_9974135.P_110.MO_1806-ID_9974135-MID_9974135-CID_30683-ST_8512-V_1)提供免费GPU时间,获取并在阿里云PAI-DSW中使用,5分钟内即可启动EasyAnimate
[![DSW Notebook](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/dsw.png)](https://gallery.pai-ml.com/#/preview/deepLearning/cv/easyanimate_v5)
#### b. 通过ComfyUI
我们的ComfyUI界面如下,具体查看[ComfyUI README](comfyui/README.md)。
![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v3/comfyui_i2v.jpg)
#### c. 通过docker
使用docker的情况下,请保证机器中已经正确安装显卡驱动与CUDA环境,然后以此执行以下命令:
```
# pull image
docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:easyanimate
# enter image
docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:easyanimate
# clone code
git clone https://github.com/aigc-apps/EasyAnimate.git
# enter EasyAnimate's dir
cd EasyAnimate
# download weights
mkdir models/Diffusion_Transformer
mkdir models/Motion_Module
mkdir models/Personalized_Model
# Please use the hugginface link or modelscope link to download the EasyAnimateV5 model.
# I2V models
# https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-InP
# https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-InP
# T2V models
# https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh
# https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh
```
### 2. 本地安装: 环境检查/下载/安装
#### a. 环境检查
我们已验证EasyAnimate可在以下环境中执行:
Windows 的详细信息:
- 操作系统 Windows 10
- python: python3.10 & python3.11
- pytorch: torch2.2.0
- CUDA: 11.8 & 12.1
- CUDNN: 8+
- GPU: Nvidia-3060 12G
Linux 的详细信息:
- 操作系统 Ubuntu 20.04, CentOS
- python: python3.10 & python3.11
- pytorch: torch2.2.0
- CUDA: 11.8 & 12.1
- CUDNN: 8+
- GPU:Nvidia-V100 16G & Nvidia-A10 24G & Nvidia-A100 40G & Nvidia-A100 80G
我们需要大约 60GB 的可用磁盘空间,请检查!
EasyAnimateV5-12B的视频大小可以由不同的GPU Memory生成,包括:
| GPU memory |384x672x25|384x672x49|576x1008x25|576x1008x49|768x1344x25|768x1344x49|
|----------|----------|----------|----------|----------|----------|----------|
| 16GB | 🧡 | 🧡 | ❌ | ❌ | ❌ | ❌ |
| 24GB | 🧡 | 🧡 | 🧡 | 🧡 | 🧡 | ❌ |
| 40GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
| 80GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
EasyAnimateV5-7B的视频大小可以由不同的GPU Memory生成,包括:
| GPU memory |384x672x25|384x672x49|576x1008x25|576x1008x49|768x1344x25|768x1344x49|
|----------|----------|----------|----------|----------|----------|----------|
| 16GB | 🧡 | 🧡 | ❌ | ❌ | ❌ | ❌ |
| 24GB | ✅ | ✅ | ✅ | 🧡 | 🧡 | ❌ |
| 40GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
| 80GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
✅ 表示它可以在"model_cpu_offload"的情况下运行,🧡代表它可以在"model_cpu_offload_and_qfloat8"的情况下运行,⭕️ 表示它可以在"sequential_cpu_offload"的情况下运行,❌ 表示它无法运行。请注意,使用sequential_cpu_offload运行会更慢。
有一些不支持torch.bfloat16的卡型,如2080ti、V100,需要将app.py、predict文件中的weight_dtype修改为torch.float16才可以运行。
EasyAnimateV5-12B使用不同GPU在25个steps中的生成时间如下:
| GPU |384x672x25|384x672x49|576x1008x25|576x1008x49|768x1344x25|768x1344x49|
|----------|----------|----------|----------|----------|----------|----------|
| A10 24GB |约120秒 (4.8s/it)|约240秒 (9.6s/it)|约320秒 (12.7s/it)| 约750秒 (29.8s/it)| ❌ | ❌ |
| A100 80GB |约45秒 (1.75s/it)|约90秒 (3.7s/it)|约120秒 (4.7s/it)|约300秒 (11.4s/it)|约265秒 (10.6s/it)| 约710秒 (28.3s/it)|
(⭕️) 表示它可以在low_gpu_memory_mode=True的情况下运行,但速度较慢,同时❌ 表示它无法运行。
<details>
<summary>(Obsolete) EasyAnimateV3:</summary>
EasyAnimateV3的视频大小可以由不同的GPU Memory生成,包括:
| GPU memory | 384x672x25 | 384x672x144 | 576x1008x72 | 576x1008x144 | 720x1280x72 | 720x1280x144 |
|----------|----------|----------|----------|----------|----------|----------|
| 12GB | ⭕️ | ⭕️ | ⭕️ | ⭕️ | ❌ | ❌ |
| 16GB | ✅ | ✅ | ⭕️ | ⭕️ | ⭕️ | ❌ |
| 24GB | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ |
| 40GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
| 80GB | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
</details>
#### b. 权重放置
我们最好将[权重](#model-zoo)按照指定路径进行放置:
EasyAnimateV5:
```
📦 models/
├── 📂 Diffusion_Transformer/
│ ├── 📂 EasyAnimateV5-12b-zh-InP/
│ └── 📂 EasyAnimateV5-12b-zh/
├── 📂 Personalized_Model/
│ └── your trained trainformer model / your trained lora model (for UI load)
```
# 视频作品
所展示的结果都是图生视频获得。
### EasyAnimateV5-12b-zh-InP
#### I2V
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/bb393b7c-ba33-494c-ab06-b314adea9fc1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/cb0d0253-919d-4dd6-9dc1-5cd94443c7f1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/09ed361f-c0c5-4025-aad7-71fe1a1a52b1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/9f42848d-34eb-473f-97ea-a5ebd0268106" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/903fda91-a0bd-48ee-bf64-fff4e4d96f17" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/407c6628-9688-44b6-b12d-77de10fbbe95" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/ccf30ec1-91d2-4d82-9ce0-fcc585fc2f21" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/5dfe0f92-7d0d-43e0-b7df-0ff7b325663c" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/2b542b85-be19-4537-9607-9d28ea7e932e" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/c1662745-752d-4ad2-92bc-fe53734347b2" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/8bec3d66-50a3-4af5-a381-be2c865825a0" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/bcec22f4-732c-446f-958c-2ebbfd8f94be" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
#### T2V
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/eccb0797-4feb-48e9-91d3-5769ce30142b" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/76b3db64-9c7a-4d38-8854-dba940240ceb" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/0b8fab66-8de7-44ff-bd43-8f701bad6bb7" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/9fbddf5f-7fcd-4cc6-9d7c-3bdf1d4ce59e" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/19c1742b-e417-45ac-97d6-8bf3a80d8e13" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/641e56c8-a3d9-489d-a3a6-42c50a9aeca1" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/2b16be76-518b-44c6-a69b-5c49d76df365" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/e7d9c0fc-136f-405c-9fab-629389e196be" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
### EasyAnimateV5-12b-zh-Control
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
<tr>
<td>
<video src="https://github.com/user-attachments/assets/53002ce2-dd18-4d4f-8135-b6f68364cabd" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/fce43c0b-81fa-4ab2-9ca7-78d786f520e6" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/b208b92c-5add-4ece-a200-3dbbe47b93c3" width="100%" controls autoplay loop></video>
</td>
<tr>
<td>
<video src="https://github.com/user-attachments/assets/3aec95d5-d240-49fb-a9e9-914446c7a4cf" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/60fa063b-5c1f-485f-b663-09bd6669de3f" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://github.com/user-attachments/assets/4adde728-8397-42f3-8a2a-23f7b39e9a1e" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
# 如何使用
<h3 id="video-gen">1. 生成 </h3>
#### a、运行python文件
- 步骤1:下载对应[权重](#model-zoo)放入models文件夹。
- 步骤2:在predict_t2v.py文件中修改prompt、neg_prompt、guidance_scale和seed。
- 步骤3:运行predict_t2v.py文件,等待生成结果,结果保存在samples/easyanimate-videos文件夹中。
- 步骤4:如果想结合自己训练的其他backbone与Lora,则看情况修改predict_t2v.py中的predict_t2v.py和lora_path。
#### b、通过ui界面
- 步骤1:下载对应[权重](#model-zoo)放入models文件夹。
- 步骤2:运行app.py文件,进入gradio页面。
- 步骤3:根据页面选择生成模型,填入prompt、neg_prompt、guidance_scale和seed等,点击生成,等待生成结果,结果保存在sample文件夹中。
#### c、通过comfyui
具体查看[ComfyUI README](comfyui/README.md)。
#### d、显存节省方案
由于EasyAnimateV5的参数非常大,我们需要考虑显存节省方案,以节省显存适应消费级显卡。我们给每个预测文件都提供了GPU_memory_mode,可以在model_cpu_offload,model_cpu_offload_and_qfloat8,sequential_cpu_offload中进行选择。
- model_cpu_offload代表整个模型在使用后会进入cpu,可以节省部分显存。
- model_cpu_offload_and_qfloat8代表整个模型在使用后会进入cpu,并且对transformer模型进行了float8的量化,可以节省更多的显存。
- sequential_cpu_offload代表模型的每一层在使用后会进入cpu,速度较慢,节省大量显存。
qfloat8会降低模型的性能,但可以节省更多的显存。如果显存足够,推荐使用model_cpu_offload。
### 2. 模型训练
一个完整的EasyAnimate训练链路应该包括数据预处理、Video VAE训练、Video DiT训练。其中Video VAE训练是一个可选项,因为我们已经提供了训练好的Video VAE。
<h4 id="data-preprocess">a.数据预处理</h4>
我们给出了两个简单的demo:
- 通过图片数据训练lora模型,详情可以查看[wiki](https://github.com/aigc-apps/EasyAnimate/wiki/Training-Lora)。
- 通过视频数据进行SFT模型,详情可以查看[wiki](https://github.com/aigc-apps/EasyAnimate/wiki/Training-SFT)。
一个完整的长视频切分、清洗、描述的数据预处理链路可以参考video caption部分的[README](easyanimate/video_caption/README.md)进行。
如果期望训练一个文生图视频的生成模型,您需要以这种格式排列数据集。
```
📦 project/
├── 📂 datasets/
│ ├── 📂 internal_datasets/
│ ├── 📂 train/
│ │ ├── 📄 00000001.mp4
│ │ ├── 📄 00000002.jpg
│ │ └── 📄 .....
│ └── 📄 json_of_internal_datasets.json
```
json_of_internal_datasets.json是一个标准的json文件。json中的file_path可以被设置为相对路径,如下所示:
```json
[
{
"file_path": "train/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
{
"file_path": "train/00000002.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
.....
]
```
你也可以将路径设置为绝对路径:
```json
[
{
"file_path": "/mnt/data/videos/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
{
"file_path": "/mnt/data/train/00000001.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
.....
]
```
<h4 id="vae-train">b. Video VAE训练 (可选)</h4>
Video VAE训练是一个可选项,因为我们已经提供了训练好的Video VAE。
如果想要进行训练,可以参考video vae部分的[README](easyanimate/vae/README.md)进行。
<h4 id="dit-train">c. Video DiT训练 </h4>
如果数据预处理时,数据的格式为相对路径,则进入scripts/train.sh进行如下设置。
```
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/json_of_internal_datasets.json"
...
train_data_format="normal"
```
如果数据的格式为绝对路径,则进入scripts/train.sh进行如下设置。
```
export DATASET_NAME=""
export DATASET_META_NAME="/mnt/data/json_of_internal_datasets.json"
```
最后运行scripts/train.sh。
```sh
sh scripts/train.sh
```
关于一些参数的设置细节,可以查看[Readme Train](scripts/README_TRAIN.md)与[Readme Lora](scripts/README_TRAIN_LORA.md)
<details>
<summary>(Obsolete) EasyAnimateV1:</summary>
如果你想训练EasyAnimateV1。请切换到git分支v1。
</details>
# 模型地址
EasyAnimateV5:
7B:
| 名称 | 种类 | 存储空间 | Hugging Face | Model Scope | 描述 |
|--|--|--|--|--|--|
| EasyAnimateV5-7b-zh-InP | EasyAnimateV5 | 22 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-7b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-7b-zh-InP)| 官方的7B图生视频权重。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以49帧、每秒8帧进行训练,支持中文与英文双语预测 |
| EasyAnimateV5-7b-zh | EasyAnimateV5 | 22 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-7b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh)| 官方的7B文生视频权重。可用于进行下游任务的fientune。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以49帧、每秒8帧进行训练,支持中文与英文双语预测 |
| EasyAnimateV5-Reward-LoRAs | EasyAnimateV5 | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-Reward-LoRAs) | 通过奖励反向传播技术,优化了EasyAnimateV5-12b生成的视频,以更好地匹配人类偏好|
12B:
| 名称 | 种类 | 存储空间 | Hugging Face | Model Scope | 描述 |
|--|--|--|--|--|--|
| EasyAnimateV5-12b-zh-InP | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-InP)| 官方的图生视频权重。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以49帧、每秒8帧进行训练,支持中文与英文双语预测 |
| EasyAnimateV5-12b-zh-Control | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-Control) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-Control)| 官方的视频控制权重,支持不同的控制条件,如Canny、Depth、Pose、MLSD等。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以49帧、每秒8帧进行训练,支持中文与英文双语预测 |
| EasyAnimateV5-12b-zh | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh)| 官方的文生视频权重。可用于进行下游任务的fientune。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以49帧、每秒8帧进行训练,支持中文与英文双语预测 |
| EasyAnimateV5-Reward-LoRAs | EasyAnimateV5 | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-Reward-LoRAs) | 通过奖励反向传播技术,优化了EasyAnimateV5-12b生成的视频,以更好地匹配人类偏好|
<details>
<summary>(Obsolete) EasyAnimateV4:</summary>
| 名称 | 种类 | 存储空间 | Hugging Face | Model Scope | 描述 |
|--|--|--|--|--|--|
| EasyAnimateV4-XL-2-InP.tar.gz | EasyAnimateV4 | 解压前 8.9 GB / 解压后 14.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV4-XL-2-InP)| [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV4-XL-2-InP)| 官方的图生视频权重。支持多分辨率(512,768,1024,1280)的视频预测,以144帧、每秒24帧进行训练 |
</details>
<details>
<summary>(Obsolete) EasyAnimateV3:</summary>
| 名称 | 种类 | 存储空间 | Hugging Face | Model Scope | 描述 |
|--|--|--|--|--|--|
| EasyAnimateV3-XL-2-InP-512x512.tar | EasyAnimateV3 | 18.2GB| [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-512x512)| [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV3-XL-2-InP-512x512)| 官方的512x512分辨率的图生视频权重。以144帧、每秒24帧进行训练 |
| EasyAnimateV3-XL-2-InP-768x768.tar | EasyAnimateV3 | 18.2GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-768x768) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV3-XL-2-InP-768x768)| 官方的768x768分辨率的图生视频权重。以144帧、每秒24帧进行训练 |
| EasyAnimateV3-XL-2-InP-960x960.tar | EasyAnimateV3 | 18.2GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-960x960) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV3-XL-2-InP-960x960)| 官方的960x960(720P)分辨率的图生视频权重。以144帧、每秒24帧进行训练 |
</details>
<details>
<summary>(Obsolete) EasyAnimateV2:</summary>
| 名称 | 种类 | 存储空间 | 下载地址 | Hugging Face | Model Scope | 描述 |
|--|--|--|--|--|--|--|
| EasyAnimateV2-XL-2-512x512.tar | EasyAnimateV2 | 16.2GB | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV2-XL-2-512x512)| [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV2-XL-2-512x512)| 官方的512x512分辨率的重量。以144帧、每秒24帧进行训练 |
| EasyAnimateV2-XL-2-768x768.tar | EasyAnimateV2 | 16.2GB | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV2-XL-2-768x768) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV2-XL-2-768x768)| 官方的768x768分辨率的重量。以144帧、每秒24帧进行训练 |
| easyanimatev2_minimalism_lora.safetensors | Lora of Pixart | 485.1MB | [Download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimatev2_minimalism_lora.safetensors)| - | - | 使用特定类型的图像进行lora训练的结果。图片可从这里[下载](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/Minimalism.zip). |
</details>
<details>
<summary>(Obsolete) EasyAnimateV1:</summary>
# Model zoo
### 1、运动权重
| 名称 | 种类 | 存储空间 | 下载地址 | 描述 |
|--|--|--|--|--|
|--|--|--|--|--|
| easyanimate_v1_mm.safetensors | Motion Module | 4.1GB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Motion_Module/easyanimate_v1_mm.safetensors) | Training with 80 frames and fps 12 |
### 2、其他权重
@@ -55,387 +463,28 @@ EasyAnimate是一个基于transformer结构的pipeline,可用于生成AI动画
| PixArt-XL-2-512x512.tar | Pixart | 11.4GB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Diffusion_Transformer/PixArt-XL-2-512x512.tar)| Pixart-Alpha official weights |
| easyanimate_portrait.safetensors | Checkpoint of Pixart | 2.3GB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait.safetensors) | Training with internal portrait datasets |
| easyanimate_portrait_lora.safetensors | Lora of Pixart | 654.0MB | [download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait_lora.safetensors)| Training with internal portrait datasets |
</details>
# 未来计划
- 支持更大规模参数量的文视频生成模型。
# 生成效果
在生成风景类animation时,采样器推荐使用DPM++和Euler A。在生成人像类animation时,采样器推荐使用Euler A和Euler。
有些时候Github无法正常显示大GIF,可以通过Download GIF下载到本地查看。
使用原始的pixart checkpoint进行预测。
| Base Models | Sampler | Seed | Resolution (h x w x f) | Prompt | GenerationResult | Download |
| ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ |
| PixArt | DPM++ | 43 | 512x512x80 | A soaring drone footage captures the majestic beauty of a coastal cliff, its red and yellow stratified rock faces rich in color and against the vibrant turquoise of the sea. Seabirds can be seen taking flight around the cliff\'s precipices. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/1-cliff.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/1-cliff.gif) |
| PixArt | DPM++ | 43 | 448x640x80 | The video captures the majestic beauty of a waterfall cascading down a cliff into a serene lake. The waterfall, with its powerful flow, is the central focus of the video. The surrounding landscape is lush and green, with trees and foliage adding to the natural beauty of the scene. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/2-waterfall.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/2-waterfall.gif) |
| PixArt | DPM++ | 43 | 704x384x80 | A vibrant scene of a snowy mountain landscape. The sky is filled with a multitude of colorful hot air balloons, each floating at different heights, creating a dynamic and lively atmosphere. The balloons are scattered across the sky, some closer to the viewer, others further away, adding depth to the scene. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/3-snowy.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/3-snowy.gif) |
| PixArt | DPM++ | 43 | 448x640x64 | The vibrant beauty of a sunflower field. The sunflowers, with their bright yellow petals and dark brown centers, are in full bloom, creating a stunning contrast against the green leaves and stems. The sunflowers are arranged in neat rows, creating a sense of order and symmetry. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/4-sunflower.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/4-sunflower.gif) |
| PixArt | DPM++ | 43 | 384x704x48 | A tranquil Vermont autumn, with leaves in vibrant colors of orange and red fluttering down a mountain stream. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/5-autumn.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/5-autumn.gif) |
| PixArt | DPM++ | 43 | 704x384x48 | A vibrant underwater scene. A group of blue fish, with yellow fins, are swimming around a coral reef. The coral reef is a mix of brown and green, providing a natural habitat for the fish. The water is a deep blue, indicating a depth of around 30 feet. The fish are swimming in a circular pattern around the coral reef, indicating a sense of motion and activity. The overall scene is a beautiful representation of marine life. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/6-underwater.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/6-underwater.gif) |
| PixArt | DPM++ | 43 | 576x448x48 | Pacific coast, carmel by the blue sea ocean and peaceful waves | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/7-coast.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/7-coast.gif) |
| PixArt | DPM++ | 43 | 576x448x80 | A snowy forest landscape with a dirt road running through it. The road is flanked by trees covered in snow, and the ground is also covered in snow. The sun is shining, creating a bright and serene atmosphere. The road appears to be empty, and there are no people or animals visible in the video. The style of the video is a natural landscape shot, with a focus on the beauty of the snowy forest and the peacefulness of the road. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/8-forest.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/8-forest.gif) |
| PixArt | DPM++ | 43 | 640x448x64 | The dynamic movement of tall, wispy grasses swaying in the wind. The sky above is filled with clouds, creating a dramatic backdrop. The sunlight pierces through the clouds, casting a warm glow on the scene. The grasses are a mix of green and brown, indicating a change in seasons. The overall style of the video is naturalistic, capturing the beauty of the landscape in a realistic manner. The focus is on the grasses and their movement, with the sky serving as a secondary element. The video does not contain any human or animal elements. |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/9-grasses.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/9-grasses.gif) |
| PixArt | DPM++ | 43 | 704x384x80 | A serene night scene in a forested area. The first frame shows a tranquil lake reflecting the star-filled sky above. The second frame reveals a beautiful sunset, casting a warm glow over the landscape. The third frame showcases the night sky, filled with stars and a vibrant Milky Way galaxy. The video is a time-lapse, capturing the transition from day to night, with the lake and forest serving as a constant backdrop. The style of the video is naturalistic, emphasizing the beauty of the night sky and the peacefulness of the forest. |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/10-night.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/10-night.gif) |
| PixArt | DPM++ | 43 | 640x448x80 | Sunset over the sea. | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/11-sunset.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/11-sunset.gif) |
使用人像checkpoint进行预测。
| Base Models | Sampler | Seed | Resolution (h x w x f) | Prompt | GenerationResult | Download |
| ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ |
| Portrait | Euler A | 43 | 448x576x80 | 1girl, 3d, black hair, brown eyes, earrings, grey background, jewelry, lips, long hair, looking at viewer, photo \\(medium\\), realistic, red lips, solo | ![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/1-check.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/1-check.gif) |
| Portrait | Euler A | 43 | 448x576x80 | 1girl, bare shoulders, blurry, brown eyes, dirty, dirty face, freckles, lips, long hair, looking at viewer, realistic, sleeveless, solo, upper body |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/2-check.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/2-check.gif) |
| Portrait | Euler A | 43 | 512x512x64 | 1girl, black hair, brown eyes, earrings, grey background, jewelry, lips, looking at viewer, mole, mole under eye, neck tattoo, nose, ponytail, realistic, shirt, simple background, solo, tattoo |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/3-check.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/3-check.gif) |
| Portrait | Euler A | 43 | 576x448x64 | 1girl, black hair, lips, looking at viewer, mole, mole under eye, mole under mouth, realistic, solo |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/5-check.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/5-check.gif) |
使用人像Lora进行预测。
| Base Models | Sampler | Seed | Resolution (h x w x f) | Prompt | GenerationResult | Download |
| ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ |
| Pixart + Lora | Euler A | 43 | 512x512x64 | 1girl, 3d, black hair, brown eyes, earrings, grey background, jewelry, lips, long hair, looking at viewer, photo \\(medium\\), realistic, red lips, solo |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/1-lora.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/1-lora.gif) |
| Pixart + Lora | Euler A | 43 | 512x512x64 | 1girl, bare shoulders, blurry, brown eyes, dirty, dirty face, freckles, lips, long hair, looking at viewer, mole, mole on breast, mole on neck, mole under eye, mole under mouth, realistic, sleeveless, solo, upper body |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/2-lora.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/2-lora.gif) |
| Pixart + Lora | Euler A | 43 | 512x512x64 | 1girl, black hair, lips, looking at viewer, mole, mole under eye, mole under mouth, realistic, solo |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/5-lora.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/5-lora.gif) |
| Pixart + Lora | Euler A | 43 | 512x512x80 | 1girl, bare shoulders, blurry, blurry background, blurry foreground, bokeh, brown eyes, christmas tree, closed mouth, collarbone, depth of field, earrings, jewelry, lips, long hair, looking at viewer, photo \\(medium\\), realistic, smile, solo |![00000001](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/low_resolution/8-lora.gif) | [Download GIF](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/8-lora.gif) |
# 快速启动
### 1. 云使用: AliyunDSW/Docker
#### a. 通过阿里云 DSW
敬请期待。
#### b. 通过docker
使用docker的情况下,请保证机器中已经正确安装显卡驱动与CUDA环境,然后以此执行以下命令:
```
# 拉取镜像
docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:easyanimate
# 进入镜像
docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:easyanimate
# clone 代码
git clone https://github.com/aigc-apps/EasyAnimate.git
# 进入EasyAnimate文件夹
cd EasyAnimate
# 下载权重
mkdir models/Diffusion_Transformer
mkdir models/Motion_Module
mkdir models/Personalized_Model
wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Motion_Module/easyanimate_v1_mm.safetensors -O models/Motion_Module/easyanimate_v1_mm.safetensors
wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait.safetensors -O models/Personalized_Model/easyanimate_portrait.safetensors
wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Personalized_Model/easyanimate_portrait_lora.safetensors -O models/Personalized_Model/easyanimate_portrait_lora.safetensors
wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Diffusion_Transformer/PixArt-XL-2-512x512.tar -O models/Diffusion_Transformer/PixArt-XL-2-512x512.tar
cd models/Diffusion_Transformer/
tar -xvf PixArt-XL-2-512x512.tar
cd ../../
```
### 2. 本地安装: 环境检查/下载/安装
#### a. 环境检查
我们已验证EasyAnimate可在以下环境中执行:
Linux 的详细信息:
- 操作系统 Ubuntu 20.04, CentOS
- python: python3.10 & python3.11
- pytorch: torch2.2.0
- CUDA: 11.8
- CUDNN: 8+
- GPU: Nvidia-A10 24G & Nvidia-A100 40G & Nvidia-A100 80G
我们需要大约 60GB 的可用磁盘空间,请检查!
#### b. 权重放置
我们最好将权重按照指定路径进行放置:
```
📦 models/
├── 📂 Diffusion_Transformer/
│ └── 📂 PixArt-XL-2-512x512/
├── 📂 Motion_Module/
│ └── 📄 easyanimate_v1_mm.safetensors
├── 📂 Motion_Module/
│ ├── 📄 easyanimate_portrait.safetensors
│ └── 📄 easyanimate_portrait_lora.safetensors
```
# 如何使用
### 1. 生成
#### a. 视频生成
##### i、运行python文件
- 步骤1:下载对应权重放入models文件夹。
- 步骤2:在predict_t2v.py文件中修改prompt、neg_prompt、guidance_scale和seed。
- 步骤3:运行predict_t2v.py文件,等待生成结果,结果保存在samples/easyanimate-videos文件夹中。
- 步骤4:如果想结合自己训练的其他backbone与Lora,则看情况修改predict_t2v.py中的predict_t2v.py和lora_path。
##### ii、通过ui界面
- 步骤1:下载对应权重放入models文件夹。
- 步骤2:运行app.py文件,进入gradio页面。
- 步骤3:根据页面选择生成模型,填入prompt、neg_prompt、guidance_scale和seed等,点击生成,等待生成结果,结果保存在sample文件夹中。
### 2. 模型训练
#### a、训练视频生成模型
##### i、基于webvid数据集
如果使用webvid数据集进行训练,则需要首先下载webvid的数据集。
您需要以这种格式排列webvid数据集。
```
📦 project/
├── 📂 datasets/
│ ├── 📂 webvid/
│ ├── 📂 videos/
│ │ ├── 📄 00000001.mp4
│ │ ├── 📄 00000002.mp4
│ │ └── 📄 .....
│ └── 📄 csv_of_webvid.csv
```
然后,进入scripts/train_t2v.sh进行设置。
```
export DATASET_NAME="datasets/webvid/videos/"
export DATASET_META_NAME="datasets/webvid/csv_of_webvid.csv"
...
train_data_format="webvid"
```
最后运行scripts/train_t2v.sh。
```sh
sh scripts/train_t2v.sh
```
##### ii、基于自建数据集
如果使用内部数据集进行训练,则需要首先格式化数据集。
您需要以这种格式排列数据集。
```
📦 project/
├── 📂 datasets/
│ ├── 📂 internal_datasets/
│ ├── 📂 videos/
│ │ ├── 📄 00000001.mp4
│ │ ├── 📄 00000002.mp4
│ │ └── 📄 .....
│ └── 📄 json_of_internal_datasets.json
```
json_of_internal_datasets.json是一个标准的json文件,如下所示:
```json
[
{
"file_path": "videos/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
{
"file_path": "videos/00000002.mp4",
"text": "A notepad with a drawing of a woman on it.",
"type": "video"
}
.....
]
```
json中的file_path需要设置为相对路径。
然后,进入scripts/train_t2v.sh进行设置。
```
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/json_of_internal_datasets.json"
...
train_data_format="normal"
```
最后运行scripts/train_t2v.sh。
```sh
sh scripts/train_t2v.sh
```
##### iii、标注视频数据集(可选)
如果需要获取更为详细且丰富的视频描述,请参考 [video_caption](./easyanimate/video_caption/)。
#### b、训练基础文生图模型
##### i、基于diffusers格式
数据集的格式可以设置为diffusers格式。
```
📦 project/
├── 📂 datasets/
│ ├── 📂 diffusers_datasets/
│ ├── 📂 train/
│ │ ├── 📄 00000001.jpg
│ │ ├── 📄 00000002.jpg
│ │ └── 📄 .....
│ └── 📄 metadata.jsonl
```
然后,进入scripts/train_t2i.sh进行设置。
```
export DATASET_NAME="datasets/diffusers_datasets/"
...
train_data_format="diffusers"
```
最后运行scripts/train_t2i.sh。
```sh
sh scripts/train_t2i.sh
```
##### ii、基于自建数据集
如果使用自建数据集进行训练,则需要首先格式化数据集。
您需要以这种格式排列数据集。
```
📦 project/
├── 📂 datasets/
│ ├── 📂 internal_datasets/
│ ├── 📂 train/
│ │ ├── 📄 00000001.jpg
│ │ ├── 📄 00000002.jpg
│ │ └── 📄 .....
│ └── 📄 json_of_internal_datasets.json
```
json_of_internal_datasets.json是一个标准的json文件,如下所示:
```json
[
{
"file_path": "train/00000001.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
{
"file_path": "train/00000002.jpg",
"text": "A notepad with a drawing of a woman on it.",
"type": "image"
}
.....
]
```
json中的file_path需要设置为相对路径。
然后,进入scripts/train_t2i.sh进行设置。
```
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/json_of_internal_datasets.json"
...
train_data_format="normal"
```
最后运行scripts/train_t2i.sh。
```sh
sh scripts/train_t2i.sh
```
#### c、训练Lora模型
##### i、基于diffusers格式
数据集的格式可以设置为diffusers格式。
```
📦 project/
├── 📂 datasets/
│ ├── 📂 diffusers_datasets/
│ ├── 📂 train/
│ │ ├── 📄 00000001.jpg
│ │ ├── 📄 00000002.jpg
│ │ └── 📄 .....
│ └── 📄 metadata.jsonl
```
然后,进入scripts/train_lora.sh进行设置。
```
export DATASET_NAME="datasets/diffusers_datasets/"
...
train_data_format="diffusers"
```
最后运行scripts/train_lora.sh。
```sh
sh scripts/train_lora.sh
```
##### ii、基于自建数据集
如果使用自建数据集进行训练,则需要首先格式化数据集。
您需要以这种格式排列数据集。
```
📦 project/
├── 📂 datasets/
│ ├── 📂 internal_datasets/
│ ├── 📂 train/
│ │ ├── 📄 00000001.jpg
│ │ ├── 📄 00000002.jpg
│ │ └── 📄 .....
│ └── 📄 json_of_internal_datasets.json
```
json_of_internal_datasets.json是一个标准的json文件,如下所示:
```json
[
{
"file_path": "train/00000001.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
{
"file_path": "train/00000002.jpg",
"text": "A notepad with a drawing of a woman on it.",
"type": "image"
}
.....
]
```
json中的file_path需要设置为相对路径。
然后,进入scripts/train_lora.sh进行设置。
```
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/json_of_internal_datasets.json"
...
train_data_format="normal"
```
最后运行scripts/train_lora.sh。
```sh
sh scripts/train_lora.sh
```
# 算法细节
我们使用了[PixArt-alpha](https://github.com/PixArt-alpha/PixArt-alpha)作为基础模型,并在此基础上引入额外的运动模块(motion module)来将DiT模型从2D图像生成扩展到3D视频生成上来。其框架图如下:
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/pipeline.png" alt="ui" style="zoom:50%;" />
其中,Motion Module 用于捕捉时序维度的帧间关系,其结构如下:
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/motion_module.png" alt="motion" style="zoom:50%;" />
我们在时序维度上引入注意力机制来让模型学习时序信息,以进行连续视频帧的生成。同时,我们利用额外的网格计算(Grid Reshape),来扩大注意力机制的input token数目,从而更多地利用图像的空间信息以达到更好的生成效果。Motion Module 作为一个单独的模块,在推理时可以用在不同的DiT基线模型上。此外,EasyAnimate不仅支持了motion-module模块的训练,也支持了DiT基模型/LoRA模型的训练,以方便用户根据自身需要来完成自定义风格的模型训练,进而生成任意风格的视频。
# 算法限制
- 受
# 联系我们
1. 扫描下方二维码或搜索群号:77450006752 来加入钉钉群。
2. 扫描下方二维码来加入微信群(如果二维码失效,可扫描最右边同学的微信,邀请您入群)
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/group/dd.png" alt="ding group" width="30%"/>
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/group/wechat.jpg" alt="Wechat group" width="30%"/>
<img src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/group/person.jpg" alt="Person" width="30%"/>
# 参考文献
- CogVideo: https://github.com/THUDM/CogVideo/
- Flux: https://github.com/black-forest-labs/flux
- magvit: https://github.com/google-research/magvit
- PixArt: https://github.com/PixArt-alpha/PixArt-alpha
- Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
- Open-Sora: https://github.com/hpcaitech/Open-Sora
- Animatediff: https://github.com/guoyww/AnimateDiff
- ComfyUI-EasyAnimateWrapper: https://github.com/kijai/ComfyUI-EasyAnimateWrapper
- HunYuan DiT: https://github.com/tencent/HunyuanDiT
# 许可证
本项目采用 [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE).
+3
View File
@@ -0,0 +1,3 @@
from .comfyui.comfyui_nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+60 -3
View File
@@ -1,6 +1,63 @@
from easyanimate.ui.ui import ui
import time
import torch
from easyanimate.api.api import (infer_forward_api,
update_diffusion_transformer_api,
update_edition_api)
from easyanimate.ui.ui import ui, ui_eas, ui_modelscope
if __name__ == "__main__":
# Choose the ui mode
ui_mode = "normal"
# GPU memory mode, which can be choosen in ["model_cpu_offload", "model_cpu_offload_and_qfloat8", "sequential_cpu_offload"].
# "model_cpu_offload" means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# "model_cpu_offload_and_qfloat8" indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# "sequential_cpu_offload" means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload"
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# Server ip
server_name = "0.0.0.0"
demo = ui()
demo.launch(server_name=server_name)
server_port = 7860
# Params below is used when ui_mode = "modelscope"
edition = "v5"
# Config
config_path = "config/easyanimate_video_v5_magvit_multi_text_encoder.yaml"
# Model path of the pretrained model
model_name = "models/Diffusion_Transformer/EasyAnimateV5-12b-zh-InP"
# "Inpaint" or "Control"
model_type = "Inpaint"
# Save dir
savedir_sample = "samples"
if ui_mode == "modelscope":
demo, controller = ui_modelscope(model_type, edition, config_path, model_name, savedir_sample, GPU_memory_mode, weight_dtype)
elif ui_mode == "eas":
demo, controller = ui_eas(edition, config_path, model_name, savedir_sample)
else:
demo, controller = ui(GPU_memory_mode, weight_dtype)
# launch gradio
app, _, _ = demo.queue(status_update_rate=1).launch(
server_name=server_name,
server_port=server_port,
prevent_thread_lock=True
)
# launch api
infer_forward_api(None, app, controller)
update_diffusion_transformer_api(None, app, controller)
update_edition_api(None, app, controller)
# not close the python
while True:
time.sleep(5)
BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 188 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 71 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 118 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 110 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 186 KiB

+101
View File
@@ -0,0 +1,101 @@
# ComfyUI EasyAnimate
Easily use EasyAnimate inside ComfyUI!
[![Arxiv Page](https://img.shields.io/badge/Arxiv-Page-red)](https://arxiv.org/abs/2405.18991)
[![Project Page](https://img.shields.io/badge/Project-Website-green)](https://easyanimate.github.io/)
[![Modelscope Studio](https://img.shields.io/badge/Modelscope-Studio-blue)](https://modelscope.cn/studios/PAI/EasyAnimate/summary)
[![Hugging Face Spaces](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Spaces-yellow)](https://huggingface.co/spaces/alibaba-pai/EasyAnimate)
- [Installation](#1-installation)
- [Node types](#node-types)
- [Example workflows](#example-workflows)
- [Image to video](#image-to-video)
- [Image to video generation (high FPS w/ frame interpolation)](#image-to-video-generation-high-fps-w-frame-interpolation)
## 1. Installation
### Option 1: Install via ComfyUI Manager
TBD
### Option 2: Install manually
The EasyAnimate repository needs to be placed at `ComfyUI/custom_nodes/EasyAnimate/`.
```
cd ComfyUI/custom_nodes/
# Git clone the easyanimate itself
git clone https://github.com/aigc-apps/EasyAnimate.git
# Git clone the video outout node
git clone https://github.com/Kosinkadink/ComfyUI-VideoHelperSuite.git
cd EasyAnimate/
pip install -r comfyui/requirements.txt
```
### 2. Download models into `ComfyUI/models/EasyAnimate/`
EasyAnimateV5:
| Name | Type | Storage Space | Hugging Face | Model Scope | Description |
|--|--|--|--|--|--|
| EasyAnimateV5-12b-zh-InP | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-InP) | Official image-to-video weights. Supports video prediction at multiple resolutions (512, 768, 1024), trained with 49 frames at 8 frames per second, and supports bilingual prediction in Chinese and English. |
| EasyAnimateV5-12b-zh-Control | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-Control) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-Control) | Official video control weights, supporting various control conditions such as Canny, Depth, Pose, MLSD, etc. Supports video prediction at multiple resolutions (512, 768, 1024) and is trained with 49 frames at 8 frames per second. Bilingual prediction in Chinese and English is supported. |
| EasyAnimateV5-12b-zh | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh) | Official text-to-video weights. Supports video prediction at multiple resolutions (512, 768, 1024), trained with 49 frames at 8 frames per second, and supports bilingual prediction in Chinese and English. |
<details>
<summary>(Obsolete) EasyAnimateV4:</summary>
| Name | Type | Storage Space | Url | Hugging Face | Description |
|--|--|--|--|--|--|
| EasyAnimateV4-XL-2-InP.tar.gz | EasyAnimateV4 | Before extraction: 8.9 GB \/ After extraction: 14.0 GB | [Download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Diffusion_Transformer/EasyAnimateV4-XL-2-InP.tar.gz) | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV4-XL-2-InP)| Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 144 frames at a rate of 24 frames per second. |
</details>
<details>
<summary>(Obsolete) EasyAnimateV3:</summary>
| Name | Type | Storage Space | Url | Hugging Face | Description |
|--|--|--|--|--|--|
| EasyAnimateV3-XL-2-InP-512x512.tar | EasyAnimateV3 | 18.2GB | [Download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Diffusion_Transformer/EasyAnimateV3-XL-2-InP-512x512.tar) | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-512x512) | EasyAnimateV3 official weights for 512x512 text and image to video resolution. Training with 144 frames and fps 24 |
| EasyAnimateV3-XL-2-InP-768x768.tar | EasyAnimateV3 | 18.2GB | [Download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Diffusion_Transformer/EasyAnimateV3-XL-2-InP-768x768.tar) | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-768x768) | EasyAnimateV3 official weights for 768x768 text and image to video resolution. Training with 144 frames and fps 24 |
| EasyAnimateV3-XL-2-InP-960x960.tar | EasyAnimateV3 | 18.2GB | [Download](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Diffusion_Transformer/EasyAnimateV3-XL-2-InP-960x960.tar) | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV3-XL-2-InP-960x960) | EasyAnimateV3 official weights for 960x960 text and image to video resolution. Training with 144 frames and fps 24 |
</details>
## Node types
- **LoadEasyAnimateModel**
- Loads the EasyAnimate model
- **EasyAnimate_TextBox**
- Write the prompt for EasyAnimate model
- **EasyAnimateI2VSampler**
- EasyAnimate Sampler for Image to Video
- **EasyAnimateT2VSampler**
- EasyAnimate Sampler for Text to Video
- **EasyAnimateV2VSampler**
- EasyAnimate Sampler for Video to Video
## Example workflows
### Video to video generation
Our ui is shown as follow, this is the [download link](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v5/easyanimatev5_workflow_v2v.json) of the json:
![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v5/easyanimatev5_workflow_v2v.jpg)
You can run the demo using following video:
[demo video](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/play_guitar.mp4)
### Control video generation
Our ui is shown as follow, this is the [download link](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v5/easyanimatev5_workflow_v2v_control.json) of the json:
![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v5/easyanimatev5_workflow_v2v_control.jpg)
You can run the demo using following video:
[demo video](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/pose.mp4)
### Image to video generation
Our ui is shown as follow, this is the [download link](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v5/easyanimatev5_workflow_i2v.json) of the json:
![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v5/easyanimatev5_workflow_i2v.jpg)
You can run the demo using following photo:
![demo image](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/firework.png)
### Text to video generation
Our ui is shown as follow, this is the [download link](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v5/easyanimatev5_workflow_t2v.json) of the json:
![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/v5/easyanimatev5_workflow_t2v.jpg)
File diff suppressed because it is too large Load Diff
+26
View File
@@ -0,0 +1,26 @@
Pillow
einops
safetensors
timm
tomesd
torch>=2.1.2
torchdiffeq
torchsde
decord
datasets
numpy
scikit-image
opencv-python
omegaconf
SentencePiece
albumentations
imageio[ffmpeg]
imageio[pyav]
tensorboard
beautifulsoup4
ftfy
func_timeout
accelerate>=0.25.0
gradio>=3.41.2,<=3.48.0
diffusers>=0.30.1
transformers>=4.37.2
+473
View File
@@ -0,0 +1,473 @@
{
"last_node_id": 81,
"last_link_id": 41,
"nodes": [
{
"id": 73,
"type": "EasyAnimate_TextBox",
"pos": [
250,
160
],
"size": {
"0": 383.7149963378906,
"1": 183.83506774902344
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
38
],
"shape": 3,
"slot_index": 0
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"The video is not of a high quality, it has a low resolution, and the audio quality is not clear. Strange motion trajectory, a poor composition and deformed video, low resolution, duplicate and ugly, strange body structure, long and strange neck, bad teeth, bad eyes, bad limbs, bad hands, rotating camera, blurry camera, shaking camera. Deformation, low-resolution, blurry, ugly, distortion."
]
},
{
"id": 7,
"type": "LoadImage",
"pos": [
258.76883544921907,
468.15773315429715
],
"size": [
378.07147216796875,
314.0000114440918
],
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
39
],
"shape": 3,
"label": "图像",
"slot_index": 0
},
{
"name": "MASK",
"type": "MASK",
"links": null,
"shape": 3,
"label": "遮罩"
}
],
"title": "Start Image(图片到视频的开始图片)",
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"firework.png",
"image"
]
},
{
"id": 79,
"type": "Note",
"pos": [
16,
460
],
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 2,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"You can upload image here\n(在此上传开始图像)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 80,
"type": "Note",
"pos": [
19.81457666015625,
-307.8177917480467
],
"size": {
"0": 210,
"1": 66.98204040527344
},
"flags": {},
"order": 3,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"Load model here\n(在此选择要使用的模型)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 78,
"type": "Note",
"pos": [
18,
-46
],
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 4,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 81,
"type": "Note",
"pos": [
789,
425
],
"size": {
"0": 248.36927795410156,
"1": 87.05973815917969
},
"flags": {},
"order": 5,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"Pay attention to selecting a base length that is compatible with the model\n(注意选择和模型相兼容的base length)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 72,
"type": "EasyAnimateI2VSampler",
"pos": [
761,
93
],
"size": {
"0": 319.20001220703125,
"1": 282
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"link": 35
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 37
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 38
},
{
"name": "start_img",
"type": "IMAGE",
"link": 39
},
{
"name": "end_img",
"type": "IMAGE",
"link": null
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
40
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "EasyAnimateI2VSampler"
},
"widgets_values": [
72,
768,
43,
"fixed",
25,
7,
"Euler"
]
},
{
"id": 17,
"type": "VHS_VideoCombine",
"pos": [
1134,
93
],
"size": [
390.9534912109375,
535.9734235491071
],
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 40,
"label": "图像",
"slot_index": 0
},
{
"name": "audio",
"type": "VHS_AUDIO",
"link": null,
"label": "音频"
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"label": "批次管理"
},
{
"name": "vae",
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"shape": 3,
"label": "文件名",
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 24,
"loop_count": 0,
"filename_prefix": "EasyAnimate",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 22,
"save_metadata": true,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "EasyAnimate_00049.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 24
}
}
}
},
{
"id": 75,
"type": "EasyAnimate_TextBox",
"pos": [
250,
-50
],
"size": {
"0": 383.54010009765625,
"1": 156.71620178222656
},
"flags": {},
"order": 6,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
37
],
"shape": 3,
"slot_index": 0
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"fireworks display over night city. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic. "
]
},
{
"id": 31,
"type": "LoadEasyAnimateModel",
"pos": [
239.81457666015626,
-307.8177917480467
],
"size": {
"0": 422.3550720214844,
"1": 154
},
"flags": {},
"order": 7,
"mode": 0,
"outputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"links": [
35
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "LoadEasyAnimateModel"
},
"widgets_values": [
"EasyAnimateV3-XL-2-InP-768x768",
"model_cpu_offload",
"Inpaint",
"easyanimate_video_v3_slicevae_motion_module.yaml",
"bf16"
]
}
],
"links": [
[
35,
31,
0,
72,
0,
"EASYANIMATESMODEL"
],
[
37,
75,
0,
72,
1,
"STRING_PROMPT"
],
[
38,
73,
0,
72,
2,
"STRING_PROMPT"
],
[
39,
7,
0,
72,
3,
"IMAGE"
],
[
40,
72,
0,
17,
0,
"IMAGE"
]
],
"groups": [
{
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24
},
{
"title": "Load EasyAnimate",
"bounding": [
220,
-388,
474,
240
],
"color": "#b06634",
"font_size": 24
},
{
"title": "Upload Your Start Image",
"bounding": [
218,
382,
452,
418
],
"color": "#a1309b",
"font_size": 24
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.7513148009015778,
"offset": [
270.9692117199942,
417.3951463953461
]
},
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
}
},
"version": 0.4
}
+383
View File
@@ -0,0 +1,383 @@
{
"last_node_id": 85,
"last_link_id": 48,
"nodes": [
{
"id": 80,
"type": "Note",
"pos": [
20,
-300
],
"size": {
"0": 210,
"1": 66.98204040527344
},
"flags": {},
"order": 0,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"Load model here\n(在此选择要使用的模型)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 78,
"type": "Note",
"pos": [
18,
-46
],
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 1,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 73,
"type": "EasyAnimate_TextBox",
"pos": [
250,
160
],
"size": {
"0": 383.7149963378906,
"1": 183.83506774902344
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
46
],
"shape": 3,
"slot_index": 0
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"The video is not of a high quality, it has a low resolution, and the audio quality is not clear. Strange motion trajectory, a poor composition and deformed video, low resolution, duplicate and ugly, strange body structure, long and strange neck, bad teeth, bad eyes, bad limbs, bad hands, rotating camera, blurry camera, shaking camera. Deformation, low-resolution, blurry, ugly, distortion."
]
},
{
"id": 17,
"type": "VHS_VideoCombine",
"pos": [
1148,
15
],
"size": [
390.9534912109375,
535.9734235491071
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 48,
"label": "图像",
"slot_index": 0
},
{
"name": "audio",
"type": "VHS_AUDIO",
"link": null,
"label": "音频"
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"label": "批次管理"
},
{
"name": "vae",
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"shape": 3,
"label": "文件名",
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 24,
"loop_count": 0,
"filename_prefix": "EasyAnimate",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 22,
"save_metadata": true,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "EasyAnimate_00050.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 24
}
}
}
},
{
"id": 75,
"type": "EasyAnimate_TextBox",
"pos": [
250,
-50
],
"size": {
"0": 383.54010009765625,
"1": 156.71620178222656
},
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
45
],
"shape": 3,
"slot_index": 0
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic."
]
},
{
"id": 85,
"type": "EasyAnimateT2VSampler",
"pos": [
769,
15
],
"size": {
"0": 315,
"1": 290
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"link": 47,
"slot_index": 0
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 45
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 46
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
48
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EasyAnimateT2VSampler"
},
"widgets_values": [
72,
1008,
576,
false,
43,
"fixed",
25,
7,
"Euler"
]
},
{
"id": 81,
"type": "Note",
"pos": [
750,
-261
],
"size": {
"0": 376.2774963378906,
"1": 224.3558807373047
},
"flags": {},
"order": 4,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"Pay attention to selecting a width and height compatible with the model;\nThe commonly used resolution for the 512x512 model is width=672 height=384;\nThe commonly used resolution for the 768x768 model is width=1008 height=576;\nThe commonly used resolution for the 1024x1024 model is width=1244 height=720;\n\n(注意选择和模型相兼容的高和宽;\n512x512模型的常用分辨率是width=672 height=384;\n768x768模型的常用分辨率是width=1008 height=576;\n1024x1024模型的常用分辨率是width=1244 height=720;)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 31,
"type": "LoadEasyAnimateModel",
"pos": [
240,
-300
],
"size": {
"0": 422.3550720214844,
"1": 154
},
"flags": {},
"order": 5,
"mode": 0,
"outputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"links": [
47
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "LoadEasyAnimateModel"
},
"widgets_values": [
"EasyAnimateV3-XL-2-InP-768x768",
"model_cpu_offload",
"Inpaint",
"easyanimate_video_v3_slicevae_motion_module.yaml",
"bf16"
]
}
],
"links": [
[
45,
75,
0,
85,
1,
"STRING_PROMPT"
],
[
46,
73,
0,
85,
2,
"STRING_PROMPT"
],
[
47,
31,
0,
85,
0,
"EASYANIMATESMODEL"
],
[
48,
85,
0,
17,
0,
"IMAGE"
]
],
"groups": [
{
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24
},
{
"title": "Load EasyAnimate",
"bounding": [
220,
-380,
472,
240
],
"color": "#b06634",
"font_size": 24
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.9090909090909092,
"offset": [
33.575633594994244,
458.8503651453463
]
},
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
}
},
"version": 0.4
}
+450
View File
@@ -0,0 +1,450 @@
{
"last_node_id": 81,
"last_link_id": 41,
"nodes": [
{
"id": 73,
"type": "EasyAnimate_TextBox",
"pos": [
250,
160
],
"size": {
"0": 383.7149963378906,
"1": 183.83506774902344
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
38
],
"shape": 3,
"slot_index": 0
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"The video is not of a high quality, it has a low resolution, and the audio quality is not clear. Strange motion trajectory, a poor composition and deformed video, low resolution, duplicate and ugly, strange body structure, long and strange neck, bad teeth, bad eyes, bad limbs, bad hands, rotating camera, blurry camera, shaking camera. Deformation, low-resolution, blurry, ugly, distortion."
]
},
{
"id": 7,
"type": "LoadImage",
"pos": [
258.76883544921907,
468.15773315429715
],
"size": [
378.07147216796875,
314.0000114440918
],
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
39
],
"shape": 3,
"label": "图像",
"slot_index": 0
},
{
"name": "MASK",
"type": "MASK",
"links": null,
"shape": 3,
"label": "遮罩"
}
],
"title": "Start Image(图片到视频的开始图片)",
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"firework.png",
"image"
]
},
{
"id": 75,
"type": "EasyAnimate_TextBox",
"pos": [
250,
-50
],
"size": {
"0": 383.54010009765625,
"1": 156.71620178222656
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
37
],
"shape": 3,
"slot_index": 0
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"fireworks display over night city. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic."
]
},
{
"id": 79,
"type": "Note",
"pos": [
16,
460
],
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 3,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"You can upload image here\n(在此上传开始图像)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 80,
"type": "Note",
"pos": [
19.024218139648436,
-329.9677812499998
],
"size": {
"0": 210,
"1": 66.98204040527344
},
"flags": {},
"order": 4,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"Load model here\n(在此选择要使用的模型)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 78,
"type": "Note",
"pos": [
18,
-46
],
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 5,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 72,
"type": "EasyAnimateI2VSampler",
"pos": [
761,
93
],
"size": {
"0": 319.20001220703125,
"1": 282
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"link": 35
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 37
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 38
},
{
"name": "start_img",
"type": "IMAGE",
"link": 39
},
{
"name": "end_img",
"type": "IMAGE",
"link": null
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
40
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "EasyAnimateI2VSampler"
},
"widgets_values": [
72,
768,
43,
"fixed",
25,
7,
"Euler"
]
},
{
"id": 31,
"type": "LoadEasyAnimateModel",
"pos": [
244.9895975341797,
-333.2615217285155
],
"size": [
440.2642531237557,
166.77216219840386
],
"flags": {},
"order": 6,
"mode": 0,
"outputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"links": [
35
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "LoadEasyAnimateModel"
},
"widgets_values": [
"EasyAnimateV4-XL-2-InP",
"model_cpu_offload",
"Inpaint",
"easyanimate_video_v4_slicevae_multi_text_encoder.yaml",
"bf16"
]
},
{
"id": 17,
"type": "VHS_VideoCombine",
"pos": [
1134,
93
],
"size": [
390.9534912109375,
535.9734235491071
],
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 40,
"label": "图像",
"slot_index": 0
},
{
"name": "audio",
"type": "VHS_AUDIO",
"link": null,
"label": "音频"
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"label": "批次管理"
},
{
"name": "vae",
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"shape": 3,
"label": "文件名",
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 24,
"loop_count": 0,
"filename_prefix": "EasyAnimate",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 22,
"save_metadata": true,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "EasyAnimate_00055.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 24
}
}
}
}
],
"links": [
[
35,
31,
0,
72,
0,
"EASYANIMATESMODEL"
],
[
37,
75,
0,
72,
1,
"STRING_PROMPT"
],
[
38,
73,
0,
72,
2,
"STRING_PROMPT"
],
[
39,
7,
0,
72,
3,
"IMAGE"
],
[
40,
72,
0,
17,
0,
"IMAGE"
]
],
"groups": [
{
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24
},
{
"title": "Load EasyAnimate",
"bounding": [
219,
-410,
492,
259
],
"color": "#b06634",
"font_size": 24
},
{
"title": "Upload Your Start Image",
"bounding": [
218,
382,
452,
418
],
"color": "#a1309b",
"font_size": 24
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.7513148009015778,
"offset": [
237.60582500124437,
450.34259561409607
]
},
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
}
},
"version": 0.4
}
+408
View File
@@ -0,0 +1,408 @@
{
"last_node_id": 96,
"last_link_id": 77,
"nodes": [
{
"id": 80,
"type": "Note",
"pos": [
19.46723632812501,
-310.0837280273436
],
"size": {
"0": 210,
"1": 66.98204040527344
},
"flags": {},
"order": 0,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"Load model here\n(在此选择要使用的模型)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 78,
"type": "Note",
"pos": [
18,
-46
],
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 1,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 73,
"type": "EasyAnimate_TextBox",
"pos": [
250,
160
],
"size": {
"0": 383.7149963378906,
"1": 183.83506774902344
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
52
],
"shape": 3,
"slot_index": 0
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"The video is not of a high quality, it has a low resolution, and the audio quality is not clear. Strange motion trajectory, a poor composition and deformed video, low resolution, duplicate and ugly, strange body structure, long and strange neck, bad teeth, bad eyes, bad limbs, bad hands, rotating camera, blurry camera, shaking camera. Deformation, low-resolution, blurry, ugly, distortion."
]
},
{
"id": 75,
"type": "EasyAnimate_TextBox",
"pos": [
250,
-50
],
"size": {
"0": 383.54010009765625,
"1": 156.71620178222656
},
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
51
],
"shape": 3,
"slot_index": 0
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body"
]
},
{
"id": 96,
"type": "LoadEasyAnimateLora",
"pos": [
716.467236328125,
-316.0837280273436
],
"size": {
"0": 461.959228515625,
"1": 128.50164794921875
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"link": 76
}
],
"outputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"links": [
77
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "LoadEasyAnimateLora"
},
"widgets_values": [
"easyanimate/easyanimatev4_minimalism_lora.safetensors",
0.7000000000000001
]
},
{
"id": 88,
"type": "EasyAnimateT2VSampler",
"pos": [
801,
35
],
"size": {
"0": 315,
"1": 290
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"link": 77
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 51
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 52,
"slot_index": 2
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
53
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EasyAnimateT2VSampler"
},
"widgets_values": [
72,
672,
384,
false,
43,
"fixed",
25,
7,
"Euler"
]
},
{
"id": 95,
"type": "LoadEasyAnimateModel",
"pos": [
243.467236328125,
-315.0837280273436
],
"size": {
"0": 436.8020324707031,
"1": 154
},
"flags": {},
"order": 4,
"mode": 0,
"outputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"links": [
76
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "LoadEasyAnimateModel"
},
"widgets_values": [
"EasyAnimateV4-XL-2-InP",
"model_cpu_offload",
"Inpaint",
"easyanimate_video_v4_slicevae_multi_text_encoder.yaml",
"bf16"
]
},
{
"id": 17,
"type": "VHS_VideoCombine",
"pos": [
1163,
34
],
"size": [
390.9534912109375,
535.9734235491071
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 53,
"label": "图像",
"slot_index": 0
},
{
"name": "audio",
"type": "VHS_AUDIO",
"link": null,
"label": "音频"
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"label": "批次管理"
},
{
"name": "vae",
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"shape": 3,
"label": "文件名",
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 24,
"loop_count": 0,
"filename_prefix": "EasyAnimate",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 22,
"save_metadata": true,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "EasyAnimate_00054.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 24
}
}
}
}
],
"links": [
[
51,
75,
0,
88,
1,
"STRING_PROMPT"
],
[
52,
73,
0,
88,
2,
"STRING_PROMPT"
],
[
53,
88,
0,
17,
0,
"IMAGE"
],
[
76,
95,
0,
96,
0,
"EASYANIMATESMODEL"
],
[
77,
96,
0,
88,
0,
"EASYANIMATESMODEL"
]
],
"groups": [
{
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24
},
{
"title": "Load EasyAnimate",
"bounding": [
219,
-390,
985,
240
],
"color": "#b06634",
"font_size": 24
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.9090909090909091,
"offset": [
38.375070223172365,
449.4001463595484
]
},
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
}
},
"version": 0.4
}
+360
View File
@@ -0,0 +1,360 @@
{
"last_node_id": 87,
"last_link_id": 49,
"nodes": [
{
"id": 80,
"type": "Note",
"pos": [
18.171542358398433,
-312.8636376953127
],
"size": {
"0": 210,
"1": 66.98204040527344
},
"flags": {},
"order": 0,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"Load model here\n(在此选择要使用的模型)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 78,
"type": "Note",
"pos": [
18,
-46
],
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 1,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 73,
"type": "EasyAnimate_TextBox",
"pos": [
250,
160
],
"size": {
"0": 383.7149963378906,
"1": 183.83506774902344
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
46
],
"shape": 3,
"slot_index": 0
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"The video is not of a high quality, it has a low resolution, and the audio quality is not clear. Strange motion trajectory, a poor composition and deformed video, low resolution, duplicate and ugly, strange body structure, long and strange neck, bad teeth, bad eyes, bad limbs, bad hands, rotating camera, blurry camera, shaking camera. Deformation, low-resolution, blurry, ugly, distortion."
]
},
{
"id": 17,
"type": "VHS_VideoCombine",
"pos": [
1148,
15
],
"size": [
390.9534912109375,
535.9734235491071
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 48,
"label": "图像",
"slot_index": 0
},
{
"name": "audio",
"type": "VHS_AUDIO",
"link": null,
"label": "音频"
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"label": "批次管理"
},
{
"name": "vae",
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"shape": 3,
"label": "文件名",
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 24,
"loop_count": 0,
"filename_prefix": "EasyAnimate",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 22,
"save_metadata": true,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "EasyAnimate_00052.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 24
}
}
}
},
{
"id": 75,
"type": "EasyAnimate_TextBox",
"pos": [
250,
-50
],
"size": {
"0": 383.54010009765625,
"1": 156.71620178222656
},
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
45
],
"shape": 3,
"slot_index": 0
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic."
]
},
{
"id": 85,
"type": "EasyAnimateT2VSampler",
"pos": [
769,
15
],
"size": {
"0": 315,
"1": 290
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"link": 49,
"slot_index": 0
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 45
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 46
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
48
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EasyAnimateT2VSampler"
},
"widgets_values": [
72,
1008,
576,
false,
43,
"fixed",
25,
7,
"Euler"
]
},
{
"id": 87,
"type": "LoadEasyAnimateModel",
"pos": [
252,
-308
],
"size": [
441.4525528088707,
154
],
"flags": {},
"order": 4,
"mode": 0,
"outputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"links": [
49
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "LoadEasyAnimateModel"
},
"widgets_values": [
"EasyAnimateV4-XL-2-InP",
"model_cpu_offload",
"Inpaint",
"easyanimate_video_v4_slicevae_multi_text_encoder.yaml",
"bf16"
]
}
],
"links": [
[
45,
75,
0,
85,
1,
"STRING_PROMPT"
],
[
46,
73,
0,
85,
2,
"STRING_PROMPT"
],
[
48,
85,
0,
17,
0,
"IMAGE"
],
[
49,
87,
0,
85,
0,
"EASYANIMATESMODEL"
]
],
"groups": [
{
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24
},
{
"title": "Load EasyAnimate",
"bounding": [
218,
-393,
503,
254
],
"color": "#b06634",
"font_size": 24
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.8264462809917354,
"offset": [
159.35166594112934,
529.544431725907
]
},
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
}
},
"version": 0.4
}
+495
View File
@@ -0,0 +1,495 @@
{
"last_node_id": 86,
"last_link_id": 47,
"nodes": [
{
"id": 80,
"type": "Note",
"pos": [
18.27766899414062,
-307.4300560058589
],
"size": {
"0": 210,
"1": 66.98204040527344
},
"flags": {},
"order": 0,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"Load model here\n(在此选择要使用的模型)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 78,
"type": "Note",
"pos": [
18,
-46
],
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 1,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 73,
"type": "EasyAnimate_TextBox",
"pos": [
250,
160
],
"size": {
"0": 383.7149963378906,
"1": 183.83506774902344
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
45
],
"shape": 3,
"slot_index": 0
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"The video is not of a high quality, it has a low resolution, and the audio quality is not clear. Strange motion trajectory, a poor composition and deformed video, low resolution, duplicate and ugly, strange body structure, long and strange neck, bad teeth, bad eyes, bad limbs, bad hands, rotating camera, blurry camera, shaking camera. Deformation, low-resolution, blurry, ugly, distortion."
]
},
{
"id": 79,
"type": "Note",
"pos": [
15.739953613281248,
462.38664912015946
],
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 3,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"You can upload video here\n(在此上传视频)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 75,
"type": "EasyAnimate_TextBox",
"pos": [
250,
-50
],
"size": {
"0": 383.54010009765625,
"1": 156.71620178222656
},
"flags": {},
"order": 4,
"mode": 0,
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
44
],
"shape": 3,
"slot_index": 0
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"一个漂亮的女人在弹吉他。视频质量高,画面清晰。高质量,杰作,最好的质量,高分辨率,超仔细。"
]
},
{
"id": 17,
"type": "VHS_VideoCombine",
"pos": [
1134,
93
],
"size": [
390.9534912109375,
535.9734235491071
],
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 47,
"label": "图像",
"slot_index": 0
},
{
"name": "audio",
"type": "VHS_AUDIO",
"link": null,
"label": "音频"
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"label": "批次管理"
},
{
"name": "vae",
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"shape": 3,
"label": "文件名",
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 24,
"loop_count": 0,
"filename_prefix": "EasyAnimate",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 22,
"save_metadata": true,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "EasyAnimate_00053.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 24
}
}
}
},
{
"id": 31,
"type": "LoadEasyAnimateModel",
"pos": [
238.27766899414075,
-307.4300560058589
],
"size": [
482.82215786889094,
154
],
"flags": {},
"order": 5,
"mode": 0,
"outputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"links": [
43
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "LoadEasyAnimateModel"
},
"widgets_values": [
"EasyAnimateV4-XL-2-InP",
"model_cpu_offload",
"Inpaint",
"easyanimate_video_v4_slicevae_multi_text_encoder.yaml",
"bf16"
]
},
{
"id": 86,
"type": "EasyAnimateV2VSampler",
"pos": [
762,
41
],
"size": {
"0": 319.20001220703125,
"1": 306
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"link": 43,
"slot_index": 0
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 44
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 45
},
{
"name": "validation_video",
"type": "IMAGE",
"link": 46,
"slot_index": 3
},
{
"name": "control_video",
"type": "IMAGE",
"link": null
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
47
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EasyAnimateV2VSampler"
},
"widgets_values": [
48,
768,
43,
"fixed",
25,
7,
0.7000000000000001,
"Euler"
]
},
{
"id": 85,
"type": "VHS_LoadVideo",
"pos": [
336,
470
],
"size": [
235.1999969482422,
398.971426827567
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
46
],
"shape": 3,
"slot_index": 0
},
{
"name": "frame_count",
"type": "INT",
"links": null,
"shape": 3
},
{
"name": "audio",
"type": "AUDIO",
"links": null,
"shape": 3
},
{
"name": "video_info",
"type": "VHS_VIDEOINFO",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "VHS_LoadVideo"
},
"widgets_values": {
"video": "1.mp4",
"force_rate": 0,
"force_size": "Disabled",
"custom_width": 512,
"custom_height": 512,
"frame_load_cap": 0,
"skip_first_frames": 0,
"select_every_nth": 1,
"choose video to upload": "image",
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"frame_load_cap": 0,
"skip_first_frames": 0,
"force_rate": 0,
"filename": "1.mp4",
"type": "input",
"format": "video/mp4",
"select_every_nth": 1
}
}
}
}
],
"links": [
[
43,
31,
0,
86,
0,
"EASYANIMATESMODEL"
],
[
44,
75,
0,
86,
1,
"STRING_PROMPT"
],
[
45,
73,
0,
86,
2,
"STRING_PROMPT"
],
[
46,
85,
0,
86,
3,
"IMAGE"
],
[
47,
86,
0,
17,
0,
"IMAGE"
]
],
"groups": [
{
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24
},
{
"title": "Load EasyAnimate",
"bounding": [
218,
-387,
542,
248
],
"color": "#b06634",
"font_size": 24
},
{
"title": "Upload Your Video",
"bounding": [
218,
385,
456,
498
],
"color": "#a1309b",
"font_size": 24
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.683013455365071,
"offset": [
284.93704305884313,
442.10948247114595
]
},
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
}
},
"version": 0.4
}
+495
View File
@@ -0,0 +1,495 @@
{
"last_node_id": 84,
"last_link_id": 48,
"nodes": [
{
"id": 79,
"type": "Note",
"pos": {
"0": 16,
"1": 460
},
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"You can upload image here\n(在此上传开始图像)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 80,
"type": "Note",
"pos": {
"0": 19.02421760559082,
"1": -329.9677734375
},
"size": {
"0": 210,
"1": 66.98204040527344
},
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"Load model here\n(在此选择要使用的模型)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 17,
"type": "VHS_VideoCombine",
"pos": {
"0": 1134,
"1": 93
},
"size": [
390.9534912109375,
310
],
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 42,
"slot_index": 0,
"label": "图像",
"shape": 7
},
{
"name": "audio",
"type": "AUDIO",
"link": null,
"label": "音频",
"shape": 7
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"label": "批次管理",
"shape": 7
},
{
"name": "vae",
"type": "VAE",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"slot_index": 0,
"shape": 3,
"label": "文件名"
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 8,
"loop_count": 0,
"filename_prefix": "EasyAnimate",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 22,
"save_metadata": true,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "EasyAnimate_00004.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 8
}
}
}
},
{
"id": 82,
"type": "EasyAnimateV5_I2VSampler",
"pos": {
"0": 767,
"1": 93
},
"size": {
"0": 336,
"1": 282
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"link": 48
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 44
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 45
},
{
"name": "start_img",
"type": "IMAGE",
"link": 47,
"shape": 7
},
{
"name": "end_img",
"type": "IMAGE",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
42
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "EasyAnimateV5_I2VSampler"
},
"widgets_values": [
49,
512,
43,
"fixed",
50,
6,
"DDIM"
]
},
{
"id": 83,
"type": "LoadEasyAnimateModel",
"pos": {
"0": 258,
"1": -324
},
"size": {
"0": 427.9729919433594,
"1": 154
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"links": [
48
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "LoadEasyAnimateModel"
},
"widgets_values": [
"EasyAnimateV5-12b-zh-InP",
"model_cpu_offload",
"Inpaint",
"easyanimate_video_v5_magvit_multi_text_encoder.yaml",
"bf16"
]
},
{
"id": 7,
"type": "LoadImage",
"pos": {
"0": 259,
"1": 468
},
"size": {
"0": 378.07147216796875,
"1": 314
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
47
],
"slot_index": 0,
"shape": 3,
"label": "图像"
},
{
"name": "MASK",
"type": "MASK",
"links": null,
"shape": 3,
"label": "遮罩"
}
],
"title": "Start Image(图片到视频的开始图片)",
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"firework.png",
"image"
]
},
{
"id": 75,
"type": "EasyAnimate_TextBox",
"pos": {
"0": 250,
"1": -50
},
"size": {
"0": 383.54010009765625,
"1": 156.71620178222656
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
44
],
"slot_index": 0,
"shape": 3
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"夜城烟花汇演"
]
},
{
"id": 73,
"type": "EasyAnimate_TextBox",
"pos": {
"0": 250,
"1": 160
},
"size": {
"0": 383.7149963378906,
"1": 183.83506774902344
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
45
],
"slot_index": 0,
"shape": 3
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"扭曲的身体,肢体残缺,文本字幕,漫画,静止,丑陋,错误,乱码。"
]
},
{
"id": 78,
"type": "Note",
"pos": {
"0": 18,
"1": -46
},
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 84,
"type": "Note",
"pos": {
"0": -98,
"1": 198
},
"size": [
326.1556114026207,
145.20905264909447
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"Using longer neg prompt such as \"Blurring, mutation, deformation, distortion, dark and solid, comics.\" can increase stability. Adding words such as \"quiet, solid\" to the neg prompt can increase dynamism.\n(使用更长的neg prompt如\"模糊,突变,变形,失真,画面暗,画面固定,连环画,漫画,线稿,没有主体。\",可以增加稳定性。在neg prompt中添加\"安静,固定\"等词语可以增加动态性。)"
],
"color": "#432",
"bgcolor": "#653"
}
],
"links": [
[
42,
82,
0,
17,
0,
"IMAGE"
],
[
44,
75,
0,
82,
1,
"STRING_PROMPT"
],
[
45,
73,
0,
82,
2,
"STRING_PROMPT"
],
[
47,
7,
0,
82,
3,
"IMAGE"
],
[
48,
83,
0,
82,
0,
"EASYANIMATESMODEL"
]
],
"groups": [
{
"title": "Upload Your Start Image",
"bounding": [
218,
382,
452,
418
],
"color": "#a1309b",
"font_size": 24,
"flags": {}
},
{
"title": "Load EasyAnimate",
"bounding": [
219,
-410,
492,
259
],
"color": "#b06634",
"font_size": 24,
"flags": {}
},
{
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24,
"flags": {}
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.7513148009015777,
"offset": [
206.91128312862938,
440.0942364134056
]
},
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
}
},
"version": 0.4
}
+400
View File
@@ -0,0 +1,400 @@
{
"last_node_id": 89,
"last_link_id": 53,
"nodes": [
{
"id": 80,
"type": "Note",
"pos": {
"0": 18.17154312133789,
"1": -312.8636474609375
},
"size": {
"0": 210,
"1": 66.98204040527344
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"Load model here\n(在此选择要使用的模型)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 78,
"type": "Note",
"pos": {
"0": 18,
"1": -46
},
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 87,
"type": "LoadEasyAnimateModel",
"pos": {
"0": 252,
"1": -308
},
"size": {
"0": 441.4525451660156,
"1": 154
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"links": [
51
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "LoadEasyAnimateModel"
},
"widgets_values": [
"EasyAnimateV5-12b-zh-InP",
"model_cpu_offload",
"Inpaint",
"easyanimate_video_v5_magvit_multi_text_encoder.yaml",
"bf16"
]
},
{
"id": 88,
"type": "EasyAnimateV5_T2VSampler",
"pos": {
"0": 786,
"1": 15
},
"size": {
"0": 327.6000061035156,
"1": 290
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"link": 51,
"slot_index": 0
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 52,
"slot_index": 1
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 53,
"slot_index": 2
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
50
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "EasyAnimateV5_T2VSampler"
},
"widgets_values": [
49,
672,
384,
false,
43,
"fixed",
50,
6,
"DDIM"
]
},
{
"id": 75,
"type": "EasyAnimate_TextBox",
"pos": {
"0": 250,
"1": -50
},
"size": {
"0": 383.54010009765625,
"1": 156.71620178222656
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
52
],
"slot_index": 0,
"shape": 3
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"一位年轻女子,有着美丽清澈的眼睛和金发,站在森林里,穿着白色的衣服,戴着皇冠。她似乎陷入了沉思,相机聚焦在她的脸上。质量高、杰作、最佳品质、高分辨率、超精细、梦幻般。"
]
},
{
"id": 73,
"type": "EasyAnimate_TextBox",
"pos": {
"0": 250,
"1": 160
},
"size": {
"0": 383.7149963378906,
"1": 183.83506774902344
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
53
],
"slot_index": 0,
"shape": 3
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"扭曲的身体,肢体残缺,文本字幕,漫画,静止,丑陋,错误,乱码。"
]
},
{
"id": 17,
"type": "VHS_VideoCombine",
"pos": {
"0": 1148,
"1": 15
},
"size": [
390.9534912109375,
310
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 50,
"slot_index": 0,
"label": "图像",
"shape": 7
},
{
"name": "audio",
"type": "AUDIO",
"link": null,
"label": "音频",
"shape": 7
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"label": "批次管理",
"shape": 7
},
{
"name": "vae",
"type": "VAE",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"slot_index": 0,
"shape": 3,
"label": "文件名"
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 8,
"loop_count": 0,
"filename_prefix": "EasyAnimate",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 22,
"save_metadata": true,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "EasyAnimate_00003.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 8
}
}
}
},
{
"id": 89,
"type": "Note",
"pos": {
"0": -97,
"1": 193
},
"size": [
326.1556114026207,
145.20905264909447
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"Using longer neg prompt such as \"Blurring, mutation, deformation, distortion, dark and solid, comics.\" can increase stability. Adding words such as \"quiet, solid\" to the neg prompt can increase dynamism.\n(使用更长的neg prompt如\"模糊,突变,变形,失真,画面暗,画面固定,连环画,漫画,线稿,没有主体。\",可以增加稳定性。在neg prompt中添加\"安静,固定\"等词语可以增加动态性。)"
],
"color": "#432",
"bgcolor": "#653"
}
],
"links": [
[
50,
88,
0,
17,
0,
"IMAGE"
],
[
51,
87,
0,
88,
0,
"EASYANIMATESMODEL"
],
[
52,
75,
0,
88,
1,
"STRING_PROMPT"
],
[
53,
73,
0,
88,
2,
"STRING_PROMPT"
]
],
"groups": [
{
"title": "Load EasyAnimate",
"bounding": [
218,
-393,
503,
254
],
"color": "#b06634",
"font_size": 24,
"flags": {}
},
{
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24,
"flags": {}
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.7513148009015777,
"offset": [
206.91128312862938,
440.0942364134056
]
},
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
}
},
"version": 0.4
}
+541
View File
@@ -0,0 +1,541 @@
{
"last_node_id": 88,
"last_link_id": 52,
"nodes": [
{
"id": 80,
"type": "Note",
"pos": {
"0": 18.27766990661621,
"1": -307.4300537109375
},
"size": {
"0": 210,
"1": 66.98204040527344
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"Load model here\n(在此选择要使用的模型)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 78,
"type": "Note",
"pos": {
"0": 18,
"1": -46
},
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 79,
"type": "Note",
"pos": {
"0": 15.739953994750977,
"1": 462.38665771484375
},
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"You can upload video here\n(在此上传视频)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 31,
"type": "LoadEasyAnimateModel",
"pos": {
"0": 238.2776641845703,
"1": -307.4300537109375
},
"size": {
"0": 482.8221435546875,
"1": 154
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"links": [
49
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "LoadEasyAnimateModel"
},
"widgets_values": [
"EasyAnimateV5-12b-zh-InP",
"model_cpu_offload",
"Inpaint",
"easyanimate_video_v5_magvit_multi_text_encoder.yaml",
"bf16"
]
},
{
"id": 73,
"type": "EasyAnimate_TextBox",
"pos": {
"0": 250,
"1": 160
},
"size": {
"0": 383.7149963378906,
"1": 183.83506774902344
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
51
],
"slot_index": 0,
"shape": 3
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"扭曲的身体,肢体残缺,文本字幕,漫画,静止,丑陋,错误,乱码。"
]
},
{
"id": 75,
"type": "EasyAnimate_TextBox",
"pos": {
"0": 250,
"1": -50
},
"size": {
"0": 383.54010009765625,
"1": 156.71620178222656
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
50
],
"slot_index": 0,
"shape": 3
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"一个漂亮的女人在弹吉他。视频质量高,画面清晰。高质量,杰作,最好的质量,高分辨率,超仔细。"
]
},
{
"id": 87,
"type": "EasyAnimateV5_V2VSampler",
"pos": {
"0": 816,
"1": 13
},
"size": {
"0": 336,
"1": 306
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"link": 49
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 50,
"slot_index": 1
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 51,
"slot_index": 2
},
{
"name": "validation_video",
"type": "IMAGE",
"link": 52,
"slot_index": 3,
"shape": 7
},
{
"name": "control_video",
"type": "IMAGE",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
48
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "EasyAnimateV5_V2VSampler"
},
"widgets_values": [
49,
512,
43,
"fixed",
35,
7,
0.7,
"DDIM"
]
},
{
"id": 17,
"type": "VHS_VideoCombine",
"pos": {
"0": 1173,
"1": 15
},
"size": [
390.9534912109375,
310
],
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 48,
"slot_index": 0,
"label": "图像",
"shape": 7
},
{
"name": "audio",
"type": "AUDIO",
"link": null,
"label": "音频",
"shape": 7
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"label": "批次管理",
"shape": 7
},
{
"name": "vae",
"type": "VAE",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"slot_index": 0,
"shape": 3,
"label": "文件名"
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 8,
"loop_count": 0,
"filename_prefix": "EasyAnimate",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 22,
"save_metadata": true,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "EasyAnimate_00005.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 8
}
}
}
},
{
"id": 85,
"type": "VHS_LoadVideo",
"pos": {
"0": 335,
"1": 476
},
"size": [
252.056640625,
262
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"shape": 7
},
{
"name": "vae",
"type": "VAE",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
52
],
"slot_index": 0,
"shape": 3
},
{
"name": "frame_count",
"type": "INT",
"links": null,
"shape": 3
},
{
"name": "audio",
"type": "AUDIO",
"links": null,
"shape": 3
},
{
"name": "video_info",
"type": "VHS_VIDEOINFO",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "VHS_LoadVideo"
},
"widgets_values": {
"video": "1.mp4",
"force_rate": 8,
"force_size": "Disabled",
"custom_width": 512,
"custom_height": 512,
"frame_load_cap": 0,
"skip_first_frames": 0,
"select_every_nth": 1,
"choose video to upload": "image",
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"frame_load_cap": 0,
"skip_first_frames": 0,
"force_rate": 8,
"filename": "1.mp4",
"type": "input",
"format": "video/mp4",
"select_every_nth": 1
}
}
}
},
{
"id": 88,
"type": "Note",
"pos": {
"0": -97,
"1": 195
},
"size": [
326.1556091308594,
145.20904541015625
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"Using longer neg prompt such as \"Blurring, mutation, deformation, distortion, dark and solid, comics.\" can increase stability. Adding words such as \"quiet, solid\" to the neg prompt can increase dynamism.\n(使用更长的neg prompt如\"模糊,突变,变形,失真,画面暗,画面固定,连环画,漫画,线稿,没有主体。\",可以增加稳定性。在neg prompt中添加\"安静,固定\"等词语可以增加动态性。)"
],
"color": "#432",
"bgcolor": "#653"
}
],
"links": [
[
48,
87,
0,
17,
0,
"IMAGE"
],
[
49,
31,
0,
87,
0,
"EASYANIMATESMODEL"
],
[
50,
75,
0,
87,
1,
"STRING_PROMPT"
],
[
51,
73,
0,
87,
2,
"STRING_PROMPT"
],
[
52,
85,
0,
87,
3,
"IMAGE"
]
],
"groups": [
{
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24,
"flags": {}
},
{
"title": "Load EasyAnimate",
"bounding": [
218,
-387,
542,
248
],
"color": "#b06634",
"font_size": 24,
"flags": {}
},
{
"title": "Upload Your Video",
"bounding": [
218,
385,
479,
529
],
"color": "#a1309b",
"font_size": 24,
"flags": {}
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.7513148009015777,
"offset": [
206.91128312862938,
440.0942364134056
]
},
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
}
},
"version": 0.4
}
@@ -0,0 +1,541 @@
{
"last_node_id": 88,
"last_link_id": 53,
"nodes": [
{
"id": 80,
"type": "Note",
"pos": {
"0": 18.27766990661621,
"1": -307.4300537109375
},
"size": {
"0": 210,
"1": 66.98204040527344
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"Load model here\n(在此选择要使用的模型)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 78,
"type": "Note",
"pos": {
"0": 18,
"1": -46
},
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 79,
"type": "Note",
"pos": {
"0": 15.739953994750977,
"1": 462.38665771484375
},
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"You can upload video here\n(在此上传视频)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 73,
"type": "EasyAnimate_TextBox",
"pos": {
"0": 250,
"1": 160
},
"size": {
"0": 383.7149963378906,
"1": 183.83506774902344
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
51
],
"slot_index": 0,
"shape": 3
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"扭曲的身体,肢体残缺,文本字幕,漫画,静止,丑陋,错误,乱码。"
]
},
{
"id": 17,
"type": "VHS_VideoCombine",
"pos": {
"0": 1173,
"1": 15
},
"size": [
390.9534912109375,
310
],
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 48,
"slot_index": 0,
"label": "图像",
"shape": 7
},
{
"name": "audio",
"type": "AUDIO",
"link": null,
"label": "音频",
"shape": 7
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"label": "批次管理",
"shape": 7
},
{
"name": "vae",
"type": "VAE",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"slot_index": 0,
"shape": 3,
"label": "文件名"
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 8,
"loop_count": 0,
"filename_prefix": "EasyAnimate",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 22,
"save_metadata": true,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "EasyAnimate_00006.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 8
}
}
}
},
{
"id": 31,
"type": "LoadEasyAnimateModel",
"pos": {
"0": 238.2776641845703,
"1": -307.4300537109375
},
"size": {
"0": 482.8221435546875,
"1": 154
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"links": [
49
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "LoadEasyAnimateModel"
},
"widgets_values": [
"EasyAnimateV5-12b-zh-Control",
"model_cpu_offload",
"Control",
"easyanimate_video_v5_magvit_multi_text_encoder.yaml",
"bf16"
]
},
{
"id": 85,
"type": "VHS_LoadVideo",
"pos": {
"0": 335,
"1": 476
},
"size": [
252.056640625,
262
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"shape": 7
},
{
"name": "vae",
"type": "VAE",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
53
],
"slot_index": 0,
"shape": 3
},
{
"name": "frame_count",
"type": "INT",
"links": null,
"shape": 3
},
{
"name": "audio",
"type": "AUDIO",
"links": null,
"shape": 3
},
{
"name": "video_info",
"type": "VHS_VIDEOINFO",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "VHS_LoadVideo"
},
"widgets_values": {
"video": "pose.mp4",
"force_rate": 0,
"force_size": "Disabled",
"custom_width": 512,
"custom_height": 512,
"frame_load_cap": 0,
"skip_first_frames": 0,
"select_every_nth": 1,
"choose video to upload": "image",
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"frame_load_cap": 0,
"skip_first_frames": 0,
"force_rate": 0,
"filename": "pose.mp4",
"type": "input",
"format": "video/mp4",
"select_every_nth": 1
}
}
}
},
{
"id": 75,
"type": "EasyAnimate_TextBox",
"pos": {
"0": 250,
"1": -50
},
"size": {
"0": 383.54010009765625,
"1": 156.71620178222656
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"links": [
50
],
"slot_index": 0,
"shape": 3
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "EasyAnimate_TextBox"
},
"widgets_values": [
"一个穿着及膝白色无袖连衣裙和白色高跟凉鞋的美女在一个光线充足、木地板的房间里跳舞。房间的背景是一扇紧闭的门、一个展示透明玻璃瓶酒精饮料的架子和一个部分可见的深色沙发。"
]
},
{
"id": 87,
"type": "EasyAnimateV5_V2VSampler",
"pos": {
"0": 816,
"1": 13
},
"size": {
"0": 336,
"1": 306
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "easyanimate_model",
"type": "EASYANIMATESMODEL",
"link": 49
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 50,
"slot_index": 1
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 51,
"slot_index": 2
},
{
"name": "validation_video",
"type": "IMAGE",
"link": null,
"slot_index": 3,
"shape": 7
},
{
"name": "control_video",
"type": "IMAGE",
"link": 53,
"shape": 7
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
48
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "EasyAnimateV5_V2VSampler"
},
"widgets_values": [
49,
512,
43,
"fixed",
35,
6,
1,
"DDIM"
]
},
{
"id": 88,
"type": "Note",
"pos": {
"0": -99,
"1": 197
},
"size": [
326.1556114026207,
145.20905264909447
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"Using longer neg prompt such as \"Blurring, mutation, deformation, distortion, dark and solid, comics.\" can increase stability. Adding words such as \"quiet, solid\" to the neg prompt can increase dynamism.\n(使用更长的neg prompt如\"模糊,突变,变形,失真,画面暗,画面固定,连环画,漫画,线稿,没有主体。\",可以增加稳定性。在neg prompt中添加\"安静,固定\"等词语可以增加动态性。)"
],
"color": "#432",
"bgcolor": "#653"
}
],
"links": [
[
48,
87,
0,
17,
0,
"IMAGE"
],
[
49,
31,
0,
87,
0,
"EASYANIMATESMODEL"
],
[
50,
75,
0,
87,
1,
"STRING_PROMPT"
],
[
51,
73,
0,
87,
2,
"STRING_PROMPT"
],
[
53,
85,
0,
87,
4,
"IMAGE"
]
],
"groups": [
{
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24,
"flags": {}
},
{
"title": "Load EasyAnimate",
"bounding": [
218,
-387,
542,
248
],
"color": "#b06634",
"font_size": 24,
"flags": {}
},
{
"title": "Upload Your Video",
"bounding": [
218,
385,
487,
789
],
"color": "#a1309b",
"font_size": 24,
"flags": {}
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.7513148009015777,
"offset": [
206.91128312862938,
440.0942364134056
]
},
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
}
},
"version": 0.4
}
-8
View File
@@ -1,8 +0,0 @@
noise_scheduler_kwargs:
beta_start: 0.0001
beta_end: 0.02
beta_schedule: "linear"
steps_offset: 1
vae_kwargs:
enable_magvit: false
@@ -1,14 +0,0 @@
transformer_additional_kwargs:
patch_3d: false
fake_3d: false
basic_block_type: "selfattentiontemporal"
time_position_encoding_before_transformer: true
noise_scheduler_kwargs:
beta_start: 0.0001
beta_end: 0.02
beta_schedule: "linear"
steps_offset: 1
vae_kwargs:
enable_magvit: false
@@ -1,4 +1,5 @@
transformer_additional_kwargs:
transformer_type: "Transformer3DModel"
patch_3d: false
fake_3d: false
basic_block_type: "motionmodule"
@@ -14,11 +15,8 @@ transformer_additional_kwargs:
temporal_attention_dim_div: 1
block_size: 2
noise_scheduler_kwargs:
beta_start: 0.0001
beta_end: 0.02
beta_schedule: "linear"
steps_offset: 1
vae_kwargs:
enable_magvit: false
vae_type: "AutoencoderKL"
text_encoder_kwargs:
enable_multi_text_encoder: false
@@ -0,0 +1,29 @@
transformer_additional_kwargs:
transformer_type: "Transformer3DModel"
patch_3d: false
fake_3d: false
basic_block_type: "motionmodule"
time_position_encoding_before_transformer: false
motion_module_type: "Vanilla"
enable_uvit: true
motion_module_kwargs:
num_attention_heads: 8
num_transformer_block: 1
attention_block_types: [ "Temporal_Self", "Temporal_Self" ]
temporal_position_encoding: true
temporal_position_encoding_max_len: 4096
temporal_attention_dim_div: 1
block_size: 1
vae_kwargs:
vae_type: "AutoencoderKLMagvit"
mini_batch_encoder: 9
mini_batch_decoder: 3
slice_mag_vae: true
slice_compression_vae: false
cache_compression_vae: false
cache_mag_vae: false
text_encoder_kwargs:
enable_multi_text_encoder: false
@@ -0,0 +1,39 @@
transformer_additional_kwargs:
transformer_type: "Transformer3DModel"
patch_3d: false
fake_3d: false
basic_block_type: "global_motionmodule"
time_position_encoding_before_transformer: false
motion_module_type: "Vanilla"
enable_uvit: true
motion_module_kwargs_even:
num_attention_heads: 16
num_transformer_block: 1
attention_block_types: [ "Temporal_Self", "Temporal_Self" ]
temporal_position_encoding: true
temporal_position_encoding_max_len: 4096
temporal_attention_dim_div: 1
block_size: 1
remove_time_embedding_in_photo: false
motion_module_kwargs_odd:
num_attention_heads: 16
num_transformer_block: 1
attention_block_types: [ "Temporal_Self", "Global_Self" ]
temporal_position_encoding: true
temporal_position_encoding_max_len: 4096
temporal_attention_dim_div: 1
block_size: 1
remove_time_embedding_in_photo: false
vae_kwargs:
vae_type: "AutoencoderKLMagvit"
mini_batch_encoder: 8
mini_batch_decoder: 2
slice_mag_vae: false
slice_compression_vae: true
cache_compression_vae: false
cache_mag_vae: false
text_encoder_kwargs:
enable_multi_text_encoder: false
@@ -0,0 +1,20 @@
transformer_additional_kwargs:
transformer_type: "HunyuanTransformer3DModel"
basic_block_type: "basic"
after_norm: false
time_position_encoding_type: "2d_rope"
time_position_encoding: true
resize_inpaint_mask_directly: false
enable_clip_in_inpaint: true
vae_kwargs:
vae_type: "AutoencoderKLMagvit"
mini_batch_encoder: 8
mini_batch_decoder: 2
slice_mag_vae: false
slice_compression_vae: false
cache_compression_vae: true
cache_mag_vae: false
text_encoder_kwargs:
enable_multi_text_encoder: true
@@ -0,0 +1,19 @@
transformer_additional_kwargs:
transformer_type: "EasyAnimateTransformer3DModel"
after_norm: false
time_position_encoding_type: "3d_rope"
resize_inpaint_mask_directly: true
enable_text_attention_mask: false
enable_clip_in_inpaint: false
vae_kwargs:
vae_type: "AutoencoderKLMagvit"
mini_batch_encoder: 4
mini_batch_decoder: 1
slice_mag_vae: false
slice_compression_vae: false
cache_compression_vae: false
cache_mag_vae: true
text_encoder_kwargs:
enable_multi_text_encoder: true
+16
View File
@@ -0,0 +1,16 @@
{
"bf16": {
"enabled": true
},
"train_micro_batch_size_per_gpu": 1,
"train_batch_size": "auto",
"gradient_accumulation_steps": "auto",
"dump_state": true,
"zero_optimization": {
"stage": 2,
"overlap_comm": true,
"contiguous_gradients": true,
"sub_group_size": 1e9,
"reduce_bucket_size": 5e8
}
}
+176
View File
@@ -0,0 +1,176 @@
import base64
import gc
import hashlib
import io
import os
import tempfile
from io import BytesIO
import gradio as gr
import torch
from fastapi import FastAPI
from PIL import Image
# Function to encode a file to Base64
def encode_file_to_base64(file_path):
with open(file_path, "rb") as file:
# Encode the data to Base64
file_base64 = base64.b64encode(file.read())
return file_base64
def update_edition_api(_: gr.Blocks, app: FastAPI, controller):
@app.post("/easyanimate/update_edition")
def _update_edition_api(
datas: dict,
):
edition = datas.get('edition', 'v2')
try:
controller.update_edition(
edition
)
comment = "Success"
except Exception as e:
torch.cuda.empty_cache()
comment = f"Error. error information is {str(e)}"
return {"message": comment}
def update_diffusion_transformer_api(_: gr.Blocks, app: FastAPI, controller):
@app.post("/easyanimate/update_diffusion_transformer")
def _update_diffusion_transformer_api(
datas: dict,
):
diffusion_transformer_path = datas.get('diffusion_transformer_path', 'none')
try:
controller.update_diffusion_transformer(
diffusion_transformer_path
)
comment = "Success"
except Exception as e:
torch.cuda.empty_cache()
comment = f"Error. error information is {str(e)}"
return {"message": comment}
def save_base64_video(base64_string):
video_data = base64.b64decode(base64_string)
md5_hash = hashlib.md5(video_data).hexdigest()
filename = f"{md5_hash}.mp4"
temp_dir = tempfile.gettempdir()
file_path = os.path.join(temp_dir, filename)
with open(file_path, 'wb') as video_file:
video_file.write(video_data)
return file_path
def save_base64_image(base64_string):
video_data = base64.b64decode(base64_string)
md5_hash = hashlib.md5(video_data).hexdigest()
filename = f"{md5_hash}.jpg"
temp_dir = tempfile.gettempdir()
file_path = os.path.join(temp_dir, filename)
with open(file_path, 'wb') as video_file:
video_file.write(video_data)
return file_path
def infer_forward_api(_: gr.Blocks, app: FastAPI, controller):
@app.post("/easyanimate/infer_forward")
def _infer_forward_api(
datas: dict,
):
base_model_path = datas.get('base_model_path', 'none')
motion_module_path = datas.get('motion_module_path', 'none')
lora_model_path = datas.get('lora_model_path', 'none')
lora_alpha_slider = datas.get('lora_alpha_slider', 0.55)
prompt_textbox = datas.get('prompt_textbox', None)
negative_prompt_textbox = datas.get('negative_prompt_textbox', 'Blurring, mutation, deformation, distortion, dark and solid, comics, text subtitles, line art.')
sampler_dropdown = datas.get('sampler_dropdown', 'Euler')
sample_step_slider = datas.get('sample_step_slider', 30)
resize_method = datas.get('resize_method', "Generate by")
width_slider = datas.get('width_slider', 672)
height_slider = datas.get('height_slider', 384)
base_resolution = datas.get('base_resolution', 512)
is_image = datas.get('is_image', False)
generation_method = datas.get('generation_method', False)
length_slider = datas.get('length_slider', 49)
overlap_video_length = datas.get('overlap_video_length', 4)
partial_video_length = datas.get('partial_video_length', 72)
cfg_scale_slider = datas.get('cfg_scale_slider', 6)
start_image = datas.get('start_image', None)
end_image = datas.get('end_image', None)
validation_video = datas.get('validation_video', None)
validation_video_mask = datas.get('validation_video_mask', None)
control_video = datas.get('control_video', None)
denoise_strength = datas.get('denoise_strength', 0.70)
seed_textbox = datas.get("seed_textbox", 43)
generation_method = "Image Generation" if is_image else generation_method
if start_image is not None:
start_image = base64.b64decode(start_image)
start_image = [Image.open(BytesIO(start_image))]
if end_image is not None:
end_image = base64.b64decode(end_image)
end_image = [Image.open(BytesIO(end_image))]
if validation_video is not None:
validation_video = save_base64_video(validation_video)
if validation_video_mask is not None:
validation_video_mask = save_base64_image(validation_video_mask)
if control_video is not None:
control_video = save_base64_video(control_video)
try:
save_sample_path, comment = controller.generate(
"",
base_model_path,
motion_module_path,
lora_model_path,
lora_alpha_slider,
prompt_textbox,
negative_prompt_textbox,
sampler_dropdown,
sample_step_slider,
resize_method,
width_slider,
height_slider,
base_resolution,
generation_method,
length_slider,
overlap_video_length,
partial_video_length,
cfg_scale_slider,
start_image,
end_image,
validation_video,
validation_video_mask,
control_video,
denoise_strength,
seed_textbox,
is_api = True,
)
except Exception as e:
gc.collect()
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
save_sample_path = ""
comment = f"Error. error information is {str(e)}"
return {"message": comment}
if save_sample_path != "":
return {"message": comment, "save_sample_path": save_sample_path, "base64_encoding": encode_file_to_base64(save_sample_path)}
else:
return {"message": comment, "save_sample_path": save_sample_path}
+95
View File
@@ -0,0 +1,95 @@
import base64
import json
import sys
import time
from datetime import datetime
from io import BytesIO
import cv2
import requests
def post_diffusion_transformer(diffusion_transformer_path, url='http://127.0.0.1:7860'):
datas = json.dumps({
"diffusion_transformer_path": diffusion_transformer_path
})
r = requests.post(f'{url}/easyanimate/update_diffusion_transformer', data=datas, timeout=1500)
data = r.content.decode('utf-8')
return data
def post_update_edition(edition, url='http://0.0.0.0:7860'):
datas = json.dumps({
"edition": edition
})
r = requests.post(f'{url}/easyanimate/update_edition', data=datas, timeout=1500)
data = r.content.decode('utf-8')
return data
def post_infer(generation_method, length_slider, url='http://127.0.0.1:7860'):
datas = json.dumps({
"base_model_path": "none",
"motion_module_path": "none",
"lora_model_path": "none",
"lora_alpha_slider": 0.55,
"prompt_textbox": "This video shows Mount saint helens, washington - the stunning scenery of a rocky mountains during golden hours - wide shot. A soaring drone footage captures the majestic beauty of a coastal cliff, its red and yellow stratified rock faces rich in color and against the vibrant turquoise of the sea.",
"negative_prompt_textbox": "Strange motion trajectory, a poor composition and deformed video, worst quality, normal quality, low quality, low resolution, duplicate and ugly, strange body structure, long and strange neck, bad teeth, bad eyes, bad limbs, bad hands, rotating camera, blurry camera, shaking camera",
"sampler_dropdown": "Euler",
"sample_step_slider": 30,
"width_slider": 672,
"height_slider": 384,
"generation_method": "Video Generation",
"length_slider": length_slider,
"cfg_scale_slider": 6,
"seed_textbox": 43,
})
r = requests.post(f'{url}/easyanimate/infer_forward', data=datas, timeout=1500)
data = r.content.decode('utf-8')
return data
if __name__ == '__main__':
# initiate time
now_date = datetime.now()
time_start = time.time()
# -------------------------- #
# Step 1: update edition
# -------------------------- #
edition = "v5"
outputs = post_update_edition(edition)
print('Output update edition: ', outputs)
# -------------------------- #
# Step 2: update edition
# -------------------------- #
diffusion_transformer_path = "models/Diffusion_Transformer/EasyAnimateV5-12b-zh-InP"
outputs = post_diffusion_transformer(diffusion_transformer_path)
print('Output update edition: ', outputs)
# -------------------------- #
# Step 3: infer
# -------------------------- #
# "Video Generation" and "Image Generation"
generation_method = "Video Generation"
length_slider = 21
outputs = post_infer(generation_method, length_slider)
# Get decoded data
outputs = json.loads(outputs)
base64_encoding = outputs["base64_encoding"]
decoded_data = base64.b64decode(base64_encoding)
is_image = True if generation_method == "Image Generation" else False
if is_image or length_slider == 1:
file_path = "1.png"
else:
file_path = "1.mp4"
with open(file_path, "wb") as file:
file.write(decoded_data)
# End of record time
# The calculated time difference is the execution time of the program, expressed in seconds / s
time_end = time.time()
time_sum = (time_end - time_start) % 60
print('# --------------------------------------------------------- #')
print(f'# Total expenditure: {time_sum}s')
print('# --------------------------------------------------------- #')
+216 -25
View File
@@ -1,8 +1,11 @@
# Copyright (c) OpenMMLab. All rights reserved.
import os
from typing import (Generic, Iterable, Iterator, List, Optional, Sequence,
Sized, TypeVar, Union)
import cv2
import numpy as np
import torch
from PIL import Image
from torch.utils.data import BatchSampler, Dataset, Sampler
@@ -43,6 +46,70 @@ def get_image_size_without_loading(path):
with Image.open(path) as img:
return img.size # (width, height)
class RandomSampler(Sampler[int]):
r"""Samples elements randomly. If without replacement, then sample from a shuffled dataset.
If with replacement, then user can specify :attr:`num_samples` to draw.
Args:
data_source (Dataset): dataset to sample from
replacement (bool): samples are drawn on-demand with replacement if ``True``, default=``False``
num_samples (int): number of samples to draw, default=`len(dataset)`.
generator (Generator): Generator used in sampling.
"""
data_source: Sized
replacement: bool
def __init__(self, data_source: Sized, replacement: bool = False,
num_samples: Optional[int] = None, generator=None) -> None:
self.data_source = data_source
self.replacement = replacement
self._num_samples = num_samples
self.generator = generator
self._pos_start = 0
if not isinstance(self.replacement, bool):
raise TypeError(f"replacement should be a boolean value, but got replacement={self.replacement}")
if not isinstance(self.num_samples, int) or self.num_samples <= 0:
raise ValueError(f"num_samples should be a positive integer value, but got num_samples={self.num_samples}")
@property
def num_samples(self) -> int:
# dataset size might change at runtime
if self._num_samples is None:
return len(self.data_source)
return self._num_samples
def __iter__(self) -> Iterator[int]:
n = len(self.data_source)
if self.generator is None:
seed = int(torch.empty((), dtype=torch.int64).random_().item())
generator = torch.Generator()
generator.manual_seed(seed)
else:
generator = self.generator
if self.replacement:
for _ in range(self.num_samples // 32):
yield from torch.randint(high=n, size=(32,), dtype=torch.int64, generator=generator).tolist()
yield from torch.randint(high=n, size=(self.num_samples % 32,), dtype=torch.int64, generator=generator).tolist()
else:
for _ in range(self.num_samples // n):
xx = torch.randperm(n, generator=generator).tolist()
if self._pos_start >= n:
self._pos_start = 0
print("xx top 10", xx[:10], self._pos_start)
for idx in range(self._pos_start, n):
yield xx[idx]
self._pos_start = (self._pos_start + 1) % n
self._pos_start = 0
yield from torch.randperm(n, generator=generator).tolist()[:self.num_samples % n]
def __len__(self) -> int:
return self.num_samples
class AspectRatioBatchImageSampler(BatchSampler):
"""A sampler wrapper for grouping images with similar aspect ratio into a same batch.
@@ -88,15 +155,21 @@ class AspectRatioBatchImageSampler(BatchSampler):
try:
image_dict = self.dataset[idx]
image_id, name = image_dict['file_path'], image_dict['text']
if self.train_folder is None:
image_dir = image_id
width, height = image_dict.get("width", None), image_dict.get("height", None)
if width is None or height is None:
image_id, name = image_dict['file_path'], image_dict['text']
if self.train_folder is None:
image_dir = image_id
else:
image_dir = os.path.join(self.train_folder, image_id)
width, height = get_image_size_without_loading(image_dir)
ratio = height / width # self.dataset[idx]
else:
image_dir = os.path.join(self.train_folder, image_id)
width, height = get_image_size_without_loading(image_dir)
ratio = height / width # self.dataset[idx]
height = int(height)
width = int(width)
ratio = height / width # self.dataset[idx]
except Exception as e:
print(e)
continue
@@ -157,24 +230,31 @@ class AspectRatioBatchSampler(BatchSampler):
for idx in self.sampler:
try:
video_dict = self.dataset[idx]
if self.train_data_format == "normal":
video_id, name = video_dict['file_path'], video_dict['text']
if self.video_folder is None:
video_dir = video_id
else:
video_dir = os.path.join(self.video_folder, video_id)
else:
videoid, name, page_dir = video_dict['videoid'], video_dict['name'], video_dict['page_dir']
video_dir = os.path.join(self.video_folder, f"{videoid}.mp4")
cap = cv2.VideoCapture(video_dir)
width, more = video_dict.get("width", None), video_dict.get("height", None)
# 获取视频尺寸
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) # 浮点数转换为整数
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) # 浮点数转换为整数
ratio = height / width # self.dataset[idx]
if width is None or height is None:
if self.train_data_format == "normal":
video_id, name = video_dict['file_path'], video_dict['text']
if self.video_folder is None:
video_dir = video_id
else:
video_dir = os.path.join(self.video_folder, video_id)
else:
videoid, name, page_dir = video_dict['videoid'], video_dict['name'], video_dict['page_dir']
video_dir = os.path.join(self.video_folder, f"{videoid}.mp4")
cap = cv2.VideoCapture(video_dir)
# 获取视频尺寸
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) # 浮点数转换为整数
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) # 浮点数转换为整数
ratio = height / width # self.dataset[idx]
else:
height = int(height)
width = int(width)
ratio = height / width # self.dataset[idx]
except Exception as e:
print(e)
print(e, self.dataset[idx], "This item is error, please check it.")
continue
# find the closest aspect ratio
closest_ratio = min(self.aspect_ratios.keys(), key=lambda r: abs(float(r) - ratio))
@@ -185,4 +265,115 @@ class AspectRatioBatchSampler(BatchSampler):
# yield a batch of indices in the same aspect ratio group
if len(bucket) == self.batch_size:
yield bucket[:]
del bucket[:]
del bucket[:]
class AspectRatioBatchImageVideoSampler(BatchSampler):
"""A sampler wrapper for grouping images with similar aspect ratio into a same batch.
Args:
sampler (Sampler): Base sampler.
dataset (Dataset): Dataset providing data information.
batch_size (int): Size of mini-batch.
drop_last (bool): If ``True``, the sampler will drop the last batch if
its size would be less than ``batch_size``.
aspect_ratios (dict): The predefined aspect ratios.
"""
def __init__(self,
sampler: Sampler,
dataset: Dataset,
batch_size: int,
train_folder: str = None,
aspect_ratios: dict = ASPECT_RATIO_512,
drop_last: bool = False
) -> None:
if not isinstance(sampler, Sampler):
raise TypeError('sampler should be an instance of ``Sampler``, '
f'but got {sampler}')
if not isinstance(batch_size, int) or batch_size <= 0:
raise ValueError('batch_size should be a positive integer value, '
f'but got batch_size={batch_size}')
self.sampler = sampler
self.dataset = dataset
self.train_folder = train_folder
self.batch_size = batch_size
self.aspect_ratios = aspect_ratios
self.drop_last = drop_last
# buckets for each aspect ratio
self.current_available_bucket_keys = list(aspect_ratios.keys())
self.bucket = {
'image':{ratio: [] for ratio in aspect_ratios},
'video':{ratio: [] for ratio in aspect_ratios}
}
def __iter__(self):
for idx in self.sampler:
content_type = self.dataset[idx].get('type', 'image')
if content_type == 'image':
try:
image_dict = self.dataset[idx]
width, height = image_dict.get("width", None), image_dict.get("height", None)
if width is None or height is None:
image_id, name = image_dict['file_path'], image_dict['text']
if self.train_folder is None:
image_dir = image_id
else:
image_dir = os.path.join(self.train_folder, image_id)
width, height = get_image_size_without_loading(image_dir)
ratio = height / width # self.dataset[idx]
else:
height = int(height)
width = int(width)
ratio = height / width # self.dataset[idx]
except Exception as e:
print(e, self.dataset[idx], "This item is error, please check it.")
continue
# find the closest aspect ratio
closest_ratio = min(self.aspect_ratios.keys(), key=lambda r: abs(float(r) - ratio))
if closest_ratio not in self.current_available_bucket_keys:
continue
bucket = self.bucket['image'][closest_ratio]
bucket.append(idx)
# yield a batch of indices in the same aspect ratio group
if len(bucket) == self.batch_size:
yield bucket[:]
del bucket[:]
else:
try:
video_dict = self.dataset[idx]
width, height = video_dict.get("width", None), video_dict.get("height", None)
if width is None or height is None:
video_id, name = video_dict['file_path'], video_dict['text']
if self.train_folder is None:
video_dir = video_id
else:
video_dir = os.path.join(self.train_folder, video_id)
cap = cv2.VideoCapture(video_dir)
# 获取视频尺寸
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) # 浮点数转换为整数
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) # 浮点数转换为整数
ratio = height / width # self.dataset[idx]
else:
height = int(height)
width = int(width)
ratio = height / width # self.dataset[idx]
except Exception as e:
print(e, self.dataset[idx], "This item is error, please check it.")
continue
# find the closest aspect ratio
closest_ratio = min(self.aspect_ratios.keys(), key=lambda r: abs(float(r) - ratio))
if closest_ratio not in self.current_available_bucket_keys:
continue
bucket = self.bucket['video'][closest_ratio]
bucket.append(idx)
# yield a batch of indices in the same aspect ratio group
if len(bucket) == self.batch_size:
yield bucket[:]
del bucket[:]
+451 -76
View File
@@ -1,9 +1,11 @@
import csv
import gc
import io
import json
import math
import os
import random
from contextlib import contextmanager
from threading import Thread
import albumentations
@@ -12,10 +14,92 @@ import numpy as np
import torch
import torchvision.transforms as transforms
from decord import VideoReader
from func_timeout import FunctionTimedOut, func_timeout
from PIL import Image
from torch.utils.data import BatchSampler, Sampler
from torch.utils.data.dataset import Dataset
VIDEO_READER_TIMEOUT = 20
def get_random_mask(shape):
f, c, h, w = shape
if f != 1:
mask_index = np.random.choice([0, 1, 2, 3, 4, 5, 6, 7, 8, 9], p=[0.05, 0.2, 0.2, 0.2, 0.05, 0.05, 0.05, 0.1, 0.05, 0.05])
else:
mask_index = np.random.choice([0, 1], p = [0.2, 0.8])
mask = torch.zeros((f, 1, h, w), dtype=torch.uint8)
if mask_index == 0:
center_x = torch.randint(0, w, (1,)).item()
center_y = torch.randint(0, h, (1,)).item()
block_size_x = torch.randint(w // 4, w // 4 * 3, (1,)).item() # 方块的宽度范围
block_size_y = torch.randint(h // 4, h // 4 * 3, (1,)).item() # 方块的高度范围
start_x = max(center_x - block_size_x // 2, 0)
end_x = min(center_x + block_size_x // 2, w)
start_y = max(center_y - block_size_y // 2, 0)
end_y = min(center_y + block_size_y // 2, h)
mask[:, :, start_y:end_y, start_x:end_x] = 1
elif mask_index == 1:
mask[:, :, :, :] = 1
elif mask_index == 2:
mask_frame_index = np.random.randint(1, 5)
mask[mask_frame_index:, :, :, :] = 1
elif mask_index == 3:
mask_frame_index = np.random.randint(1, 5)
mask[mask_frame_index:-mask_frame_index, :, :, :] = 1
elif mask_index == 4:
center_x = torch.randint(0, w, (1,)).item()
center_y = torch.randint(0, h, (1,)).item()
block_size_x = torch.randint(w // 4, w // 4 * 3, (1,)).item() # 方块的宽度范围
block_size_y = torch.randint(h // 4, h // 4 * 3, (1,)).item() # 方块的高度范围
start_x = max(center_x - block_size_x // 2, 0)
end_x = min(center_x + block_size_x // 2, w)
start_y = max(center_y - block_size_y // 2, 0)
end_y = min(center_y + block_size_y // 2, h)
mask_frame_before = np.random.randint(0, f // 2)
mask_frame_after = np.random.randint(f // 2, f)
mask[mask_frame_before:mask_frame_after, :, start_y:end_y, start_x:end_x] = 1
elif mask_index == 5:
mask = torch.randint(0, 2, (f, 1, h, w), dtype=torch.uint8)
elif mask_index == 6:
num_frames_to_mask = random.randint(1, max(f // 2, 1))
frames_to_mask = random.sample(range(f), num_frames_to_mask)
for i in frames_to_mask:
block_height = random.randint(1, h // 4)
block_width = random.randint(1, w // 4)
top_left_y = random.randint(0, h - block_height)
top_left_x = random.randint(0, w - block_width)
mask[i, 0, top_left_y:top_left_y + block_height, top_left_x:top_left_x + block_width] = 1
elif mask_index == 7:
center_x = torch.randint(0, w, (1,)).item()
center_y = torch.randint(0, h, (1,)).item()
a = torch.randint(min(w, h) // 8, min(w, h) // 4, (1,)).item() # 长半轴
b = torch.randint(min(h, w) // 8, min(h, w) // 4, (1,)).item() # 短半轴
for i in range(h):
for j in range(w):
if ((i - center_y) ** 2) / (b ** 2) + ((j - center_x) ** 2) / (a ** 2) < 1:
mask[:, :, i, j] = 1
elif mask_index == 8:
center_x = torch.randint(0, w, (1,)).item()
center_y = torch.randint(0, h, (1,)).item()
radius = torch.randint(min(h, w) // 8, min(h, w) // 4, (1,)).item()
for i in range(h):
for j in range(w):
if (i - center_y) ** 2 + (j - center_x) ** 2 < radius ** 2:
mask[:, :, i, j] = 1
elif mask_index == 9:
for idx in range(f):
if np.random.rand() > 0.5:
mask[idx, :, :, :] = 1
else:
raise ValueError(f"The mask_index {mask_index} is not define")
return mask
class ImageVideoSampler(BatchSampler):
"""A sampler wrapper for grouping images with similar aspect ratio into a same batch.
@@ -64,17 +148,48 @@ class ImageVideoSampler(BatchSampler):
yield bucket[:]
del bucket[:]
@contextmanager
def VideoReader_contextmanager(*args, **kwargs):
vr = VideoReader(*args, **kwargs)
try:
yield vr
finally:
del vr
gc.collect()
def get_video_reader_batch(video_reader, batch_index):
frames = video_reader.get_batch(batch_index).asnumpy()
return frames
def resize_frame(frame, target_short_side):
h, w, _ = frame.shape
if h < w:
if target_short_side > h:
return frame
new_h = target_short_side
new_w = int(target_short_side * w / h)
else:
if target_short_side > w:
return frame
new_w = target_short_side
new_h = int(target_short_side * h / w)
resized_frame = cv2.resize(frame, (new_w, new_h))
return resized_frame
class ImageVideoDataset(Dataset):
def __init__(
self,
ann_path, data_root=None,
video_sample_size=512, video_sample_stride=4, video_sample_n_frames=16,
image_sample_size=512,
# For Random Crop
min_crop_f=0.9, max_crop_f=1,
video_repeat=0,
enable_bucket=False
):
self,
ann_path, data_root=None,
video_sample_size=512, video_sample_stride=4, video_sample_n_frames=16,
image_sample_size=512,
video_repeat=0,
text_drop_ratio=-1,
enable_bucket=False,
video_length_drop_start=0.1,
video_length_drop_end=0.9,
enable_inpaint=False,
):
# Loading annotations from files
print(f"loading annotations from {ann_path} ...")
if ann_path.endswith('.csv'):
@@ -88,123 +203,383 @@ class ImageVideoDataset(Dataset):
# It's used to balance num of images and videos.
self.dataset = []
for data in dataset:
if data.get('data_type', 'image') != 'video' or data.get('type', 'image') != 'video':
if data.get('type', 'image') != 'video':
self.dataset.append(data)
if video_repeat > 0:
for _ in range(video_repeat):
for data in dataset:
if data.get('data_type', 'image') == 'video' or data.get('type', 'image') == 'video':
if data.get('type', 'image') == 'video':
self.dataset.append(data)
del dataset
self.length = len(self.dataset)
print(f"data scale: {self.length}")
self.min_crop_f = min_crop_f
self.max_crop_f = max_crop_f
# TODO: enable bucket training
self.enable_bucket = enable_bucket
self.text_drop_ratio = text_drop_ratio
self.enable_inpaint = enable_inpaint
self.video_length_drop_start = video_length_drop_start
self.video_length_drop_end = video_length_drop_end
# Video params
self.video_sample_stride = video_sample_stride
self.video_sample_n_frames = video_sample_n_frames
self.video_sample_size = tuple(video_sample_size) if not isinstance(video_sample_size, int) else (video_sample_size, video_sample_size)
self.video_rescaler = albumentations.SmallestMaxSize(max_size=min(self.video_sample_size), interpolation=cv2.INTER_AREA)
self.video_sample_size = tuple(video_sample_size) if not isinstance(video_sample_size, int) else (video_sample_size, video_sample_size)
self.video_transforms = transforms.Compose(
[
transforms.Resize(min(self.video_sample_size)),
transforms.CenterCrop(self.video_sample_size),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
]
)
# Image params
self.image_sample_size = tuple(image_sample_size) if not isinstance(image_sample_size, int) else (image_sample_size, image_sample_size)
self.image_transforms = transforms.Compose([
transforms.RandomHorizontalFlip(),
transforms.Resize(min(self.image_sample_size)),
transforms.CenterCrop(self.image_sample_size),
transforms.ToTensor(),
transforms.Normalize([0.5, 0.5, 0.5],[0.5, 0.5, 0.5])
])
self.larger_side_of_image_and_video = max(min(self.image_sample_size), min(self.video_sample_size))
def get_batch(self, idx):
data_info = self.dataset[idx % len(self.dataset)]
if data_info.get('data_type', 'image')=='video' or data_info.get('type', 'image')=='video':
video_path, text = data_info['file_path'], data_info['text']
# Get abs path of video
if self.data_root is not None:
video_path = os.path.join(self.data_root, video_path)
if data_info.get('type', 'image')=='video':
video_id, text = data_info['file_path'], data_info['text']
# Get video information firstly
video_reader = VideoReader(video_path, num_threads=2)
h, w, c = video_reader[0].shape
del video_reader
# Resize to bigger firstly
t_h = int(self.video_sample_size[0] * 1.25 * h / min(h, w))
t_w = int(self.video_sample_size[0] * 1.25 * w / min(h, w))
if self.data_root is None:
video_dir = video_id
else:
video_dir = os.path.join(self.data_root, video_id)
# Get video pixels
video_reader = VideoReader(video_path, width=t_w, height=t_h, num_threads=2)
video_length = len(video_reader)
clip_length = min(video_length, (self.video_sample_n_frames - 1) * self.video_sample_stride + 1)
start_idx = random.randint(0, video_length - clip_length)
batch_index = np.linspace(start_idx, start_idx + clip_length - 1, self.video_sample_n_frames, dtype=int)
imgs = video_reader.get_batch(batch_index).asnumpy()
del video_reader
if imgs.shape[0] != self.video_sample_n_frames:
raise ValueError('Video data Sampler Error')
with VideoReader_contextmanager(video_dir, num_threads=2) as video_reader:
min_sample_n_frames = min(
self.video_sample_n_frames,
int(len(video_reader) * (self.video_length_drop_end - self.video_length_drop_start) // self.video_sample_stride)
)
if min_sample_n_frames == 0:
raise ValueError(f"No Frames in video.")
# Crop center of above videos
min_side_len = min(imgs[0].shape[:2])
crop_side_len = min_side_len * np.random.uniform(self.min_crop_f, self.max_crop_f, size=None)
crop_side_len = int(crop_side_len)
self.cropper = albumentations.RandomCrop(height=crop_side_len, width=crop_side_len)
imgs = np.transpose(imgs, (1, 2, 3, 0))
imgs = self.cropper(image=imgs)["image"]
imgs = np.transpose(imgs, (3, 0, 1, 2))
out_imgs = []
video_length = int(self.video_length_drop_end * len(video_reader))
clip_length = min(video_length, (min_sample_n_frames - 1) * self.video_sample_stride + 1)
start_idx = random.randint(int(self.video_length_drop_start * video_length), video_length - clip_length) if video_length != clip_length else 0
batch_index = np.linspace(start_idx, start_idx + clip_length - 1, min_sample_n_frames, dtype=int)
# Resize to video_sample_size
for img in imgs:
img = self.video_rescaler(image=img)["image"]
out_imgs.append(img[None, :, :, :])
imgs = np.concatenate(out_imgs).transpose(0, 3, 1, 2)
try:
sample_args = (video_reader, batch_index)
pixel_values = func_timeout(
VIDEO_READER_TIMEOUT, get_video_reader_batch, args=sample_args
)
resized_frames = []
for i in range(len(pixel_values)):
frame = pixel_values[i]
resized_frame = resize_frame(frame, self.larger_side_of_image_and_video)
resized_frames.append(resized_frame)
pixel_values = np.array(resized_frames)
except FunctionTimedOut:
raise ValueError(f"Read {idx} timeout.")
except Exception as e:
raise ValueError(f"Failed to extract frames from video. Error is {e}.")
# Normalize to -1~1
imgs = ((imgs - 127.5) / 127.5).astype(np.float32)
if imgs.shape[0] != self.video_sample_n_frames:
raise ValueError('video data sampler error')
# Random use no text generation
if random.random() < 0.1:
text = ''
return torch.from_numpy(imgs), text, 'video'
if not self.enable_bucket:
pixel_values = torch.from_numpy(pixel_values).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
del video_reader
else:
pixel_values = pixel_values
if not self.enable_bucket:
pixel_values = self.video_transforms(pixel_values)
# Random use no text generation
if random.random() < self.text_drop_ratio:
text = ''
return pixel_values, text, 'video'
else:
image_path, text = data_info['file_path'], data_info['text']
if self.data_root is not None:
image_path = os.path.join(self.data_root, image_path)
image = Image.open(image_path).convert('RGB')
image = self.image_transforms(image).unsqueeze(0)
if random.random()<0.1:
if not self.enable_bucket:
image = self.image_transforms(image).unsqueeze(0)
else:
image = np.expand_dims(np.array(image), 0)
if random.random() < self.text_drop_ratio:
text = ''
return image, text, 'video'
return image, text, 'image'
def __len__(self):
return self.length
def __getitem__(self, idx):
data_info = self.dataset[idx % len(self.dataset)]
data_type = data_info.get('type', 'image')
while True:
sample = {}
def get_data(data_idx):
try:
data_info_local = self.dataset[idx % len(self.dataset)]
data_type_local = data_info_local.get('type', 'image')
if data_type_local != data_type:
raise ValueError("data_type_local != data_type")
pixel_values, name, data_type = self.get_batch(idx)
sample["pixel_values"] = pixel_values
sample["text"] = name
sample["data_type"] = data_type
sample["idx"] = idx
try:
t = Thread(target=get_data, args=(idx, ))
t.start()
t.join(5)
if len(sample)>0:
if len(sample) > 0:
break
except Exception as e:
print(self.dataset[idx])
idx = idx - 1
print(e, self.dataset[idx % len(self.dataset)])
idx = random.randint(0, self.length-1)
if self.enable_inpaint and not self.enable_bucket:
mask = get_random_mask(pixel_values.size())
mask_pixel_values = pixel_values * (1 - mask) + torch.ones_like(pixel_values) * -1 * mask
sample["mask_pixel_values"] = mask_pixel_values
sample["mask"] = mask
clip_pixel_values = sample["pixel_values"][0].permute(1, 2, 0).contiguous()
clip_pixel_values = (clip_pixel_values * 0.5 + 0.5) * 255
sample["clip_pixel_values"] = clip_pixel_values
ref_pixel_values = sample["pixel_values"][0].unsqueeze(0)
if (mask == 1).all():
ref_pixel_values = torch.ones_like(ref_pixel_values) * -1
sample["ref_pixel_values"] = ref_pixel_values
return sample
class ImageVideoControlDataset(Dataset):
def __init__(
self,
ann_path, data_root=None,
video_sample_size=512, video_sample_stride=4, video_sample_n_frames=16,
image_sample_size=512,
video_repeat=0,
text_drop_ratio=-1,
enable_bucket=False,
video_length_drop_start=0.1,
video_length_drop_end=0.9,
enable_inpaint=False,
):
# Loading annotations from files
print(f"loading annotations from {ann_path} ...")
if ann_path.endswith('.csv'):
with open(ann_path, 'r') as csvfile:
dataset = list(csv.DictReader(csvfile))
elif ann_path.endswith('.json'):
dataset = json.load(open(ann_path))
self.data_root = data_root
# It's used to balance num of images and videos.
self.dataset = []
for data in dataset:
if data.get('type', 'image') != 'video':
self.dataset.append(data)
if video_repeat > 0:
for _ in range(video_repeat):
for data in dataset:
if data.get('type', 'image') == 'video':
self.dataset.append(data)
del dataset
self.length = len(self.dataset)
print(f"data scale: {self.length}")
# TODO: enable bucket training
self.enable_bucket = enable_bucket
self.text_drop_ratio = text_drop_ratio
self.enable_inpaint = enable_inpaint
self.video_length_drop_start = video_length_drop_start
self.video_length_drop_end = video_length_drop_end
# Video params
self.video_sample_stride = video_sample_stride
self.video_sample_n_frames = video_sample_n_frames
self.video_sample_size = tuple(video_sample_size) if not isinstance(video_sample_size, int) else (video_sample_size, video_sample_size)
self.video_transforms = transforms.Compose(
[
transforms.Resize(min(self.video_sample_size)),
transforms.CenterCrop(self.video_sample_size),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
]
)
# Image params
self.image_sample_size = tuple(image_sample_size) if not isinstance(image_sample_size, int) else (image_sample_size, image_sample_size)
self.image_transforms = transforms.Compose([
transforms.Resize(min(self.image_sample_size)),
transforms.CenterCrop(self.image_sample_size),
transforms.ToTensor(),
transforms.Normalize([0.5, 0.5, 0.5],[0.5, 0.5, 0.5])
])
self.larger_side_of_image_and_video = max(min(self.image_sample_size), min(self.video_sample_size))
def get_batch(self, idx):
data_info = self.dataset[idx % len(self.dataset)]
video_id, text = data_info['file_path'], data_info['text']
if data_info.get('type', 'image')=='video':
if self.data_root is None:
video_dir = video_id
else:
video_dir = os.path.join(self.data_root, video_id)
with VideoReader_contextmanager(video_dir, num_threads=2) as video_reader:
min_sample_n_frames = min(
self.video_sample_n_frames,
int(len(video_reader) * (self.video_length_drop_end - self.video_length_drop_start) // self.video_sample_stride)
)
if min_sample_n_frames == 0:
raise ValueError(f"No Frames in video.")
video_length = int(self.video_length_drop_end * len(video_reader))
clip_length = min(video_length, (min_sample_n_frames - 1) * self.video_sample_stride + 1)
start_idx = random.randint(int(self.video_length_drop_start * video_length), video_length - clip_length) if video_length != clip_length else 0
batch_index = np.linspace(start_idx, start_idx + clip_length - 1, min_sample_n_frames, dtype=int)
try:
sample_args = (video_reader, batch_index)
pixel_values = func_timeout(
VIDEO_READER_TIMEOUT, get_video_reader_batch, args=sample_args
)
resized_frames = []
for i in range(len(pixel_values)):
frame = pixel_values[i]
resized_frame = resize_frame(frame, self.larger_side_of_image_and_video)
resized_frames.append(resized_frame)
pixel_values = np.array(resized_frames)
except FunctionTimedOut:
raise ValueError(f"Read {idx} timeout.")
except Exception as e:
raise ValueError(f"Failed to extract frames from video. Error is {e}.")
if not self.enable_bucket:
pixel_values = torch.from_numpy(pixel_values).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
del video_reader
else:
pixel_values = pixel_values
if not self.enable_bucket:
pixel_values = self.video_transforms(pixel_values)
# Random use no text generation
if random.random() < self.text_drop_ratio:
text = ''
control_video_id = data_info['control_file_path']
if self.data_root is None:
control_video_id = control_video_id
else:
control_video_id = os.path.join(self.data_root, control_video_id)
with VideoReader_contextmanager(control_video_id, num_threads=2) as control_video_reader:
try:
sample_args = (control_video_reader, batch_index)
control_pixel_values = func_timeout(
VIDEO_READER_TIMEOUT, get_video_reader_batch, args=sample_args
)
resized_frames = []
for i in range(len(control_pixel_values)):
frame = control_pixel_values[i]
resized_frame = resize_frame(frame, self.larger_side_of_image_and_video)
resized_frames.append(resized_frame)
control_pixel_values = np.array(resized_frames)
except FunctionTimedOut:
raise ValueError(f"Read {idx} timeout.")
except Exception as e:
raise ValueError(f"Failed to extract frames from video. Error is {e}.")
if not self.enable_bucket:
control_pixel_values = torch.from_numpy(control_pixel_values).permute(0, 3, 1, 2).contiguous()
control_pixel_values = control_pixel_values / 255.
del control_video_reader
else:
control_pixel_values = control_pixel_values
if not self.enable_bucket:
control_pixel_values = self.video_transforms(control_pixel_values)
return pixel_values, control_pixel_values, text, "video"
else:
image_path, text = data_info['file_path'], data_info['text']
if self.data_root is not None:
image_path = os.path.join(self.data_root, image_path)
image = Image.open(image_path).convert('RGB')
if not self.enable_bucket:
image = self.image_transforms(image).unsqueeze(0)
else:
image = np.expand_dims(np.array(image), 0)
if random.random() < self.text_drop_ratio:
text = ''
control_image_id = data_info['control_file_path']
if self.data_root is None:
control_image_id = control_image_id
else:
control_image_id = os.path.join(self.data_root, control_image_id)
control_image = Image.open(control_image_id).convert('RGB')
if not self.enable_bucket:
control_image = self.image_transforms(control_image).unsqueeze(0)
else:
control_image = np.expand_dims(np.array(control_image), 0)
return image, control_image, text, 'image'
def __len__(self):
return self.length
def __getitem__(self, idx):
data_info = self.dataset[idx % len(self.dataset)]
data_type = data_info.get('type', 'image')
while True:
sample = {}
try:
data_info_local = self.dataset[idx % len(self.dataset)]
data_type_local = data_info_local.get('type', 'image')
if data_type_local != data_type:
raise ValueError("data_type_local != data_type")
pixel_values, control_pixel_values, name, data_type = self.get_batch(idx)
sample["pixel_values"] = pixel_values
sample["control_pixel_values"] = control_pixel_values
sample["text"] = name
sample["data_type"] = data_type
sample["idx"] = idx
if len(sample) > 0:
break
except Exception as e:
print(e, self.dataset[idx % len(self.dataset)])
idx = random.randint(0, self.length-1)
if self.enable_inpaint and not self.enable_bucket:
mask = get_random_mask(pixel_values.size())
mask_pixel_values = pixel_values * (1 - mask) + torch.ones_like(pixel_values) * -1 * mask
sample["mask_pixel_values"] = mask_pixel_values
sample["mask"] = mask
clip_pixel_values = sample["pixel_values"][0].permute(1, 2, 0).contiguous()
clip_pixel_values = (clip_pixel_values * 0.5 + 0.5) * 255
sample["clip_pixel_values"] = clip_pixel_values
ref_pixel_values = sample["pixel_values"][0].unsqueeze(0)
if (mask == 1).all():
ref_pixel_values = torch.ones_like(ref_pixel_values) * -1
sample["ref_pixel_values"] = ref_pixel_values
return sample
if __name__ == "__main__":
@@ -213,4 +588,4 @@ if __name__ == "__main__":
)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=4, num_workers=16)
for idx, batch in enumerate(dataloader):
print(batch["pixel_values"].shape, len(batch["text"]))
print(batch["pixel_values"].shape, len(batch["text"]))
+48 -13
View File
@@ -1,17 +1,26 @@
import csv
import gc
import io
import json
import math
import os
import random
from contextlib import contextmanager
from threading import Thread
import albumentations
import cv2
import numpy as np
import torch
import torchvision.transforms as transforms
from decord import VideoReader
from einops import rearrange
from func_timeout import FunctionTimedOut, func_timeout
from PIL import Image
from torch.utils.data import BatchSampler, Sampler
from torch.utils.data.dataset import Dataset
VIDEO_READER_TIMEOUT = 20
def get_random_mask(shape):
f, c, h, w = shape
@@ -53,6 +62,21 @@ def get_random_mask(shape):
return mask
@contextmanager
def VideoReader_contextmanager(*args, **kwargs):
vr = VideoReader(*args, **kwargs)
try:
yield vr
finally:
del vr
gc.collect()
def get_video_reader_batch(video_reader, batch_index):
frames = video_reader.get_batch(batch_index).asnumpy()
return frames
class WebVid10M(Dataset):
def __init__(
self,
@@ -165,21 +189,32 @@ class VideoDataset(Dataset):
video_dir = video_id
else:
video_dir = os.path.join(self.video_folder, video_id)
video_reader = VideoReader(video_dir)
video_length = len(video_reader)
clip_length = min(video_length, (self.sample_n_frames - 1) * self.sample_stride + 1)
start_idx = random.randint(0, video_length - clip_length)
batch_index = np.linspace(start_idx, start_idx + clip_length - 1, self.sample_n_frames, dtype=int)
if not self.enable_bucket:
pixel_values = torch.from_numpy(video_reader.get_batch(batch_index).asnumpy()).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
del video_reader
else:
pixel_values = video_reader.get_batch(batch_index).asnumpy()
with VideoReader_contextmanager(video_dir, num_threads=2) as video_reader:
video_length = len(video_reader)
clip_length = min(video_length, (self.sample_n_frames - 1) * self.sample_stride + 1)
start_idx = random.randint(0, video_length - clip_length)
batch_index = np.linspace(start_idx, start_idx + clip_length - 1, self.sample_n_frames, dtype=int)
return pixel_values, name
try:
sample_args = (video_reader, batch_index)
pixel_values = func_timeout(
VIDEO_READER_TIMEOUT, get_video_reader_batch, args=sample_args
)
except FunctionTimedOut:
raise ValueError(f"Read {idx} timeout.")
except Exception as e:
raise ValueError(f"Failed to extract frames from video. Error is {e}.")
if not self.enable_bucket:
pixel_values = torch.from_numpy(pixel_values).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
del video_reader
else:
pixel_values = pixel_values
return pixel_values, name
def __len__(self):
return self.length
+16
View File
@@ -0,0 +1,16 @@
from .autoencoder_magvit import (AutoencoderKLCogVideoX, AutoencoderKLMagvit, AutoencoderKL)
from .transformer3d import (EasyAnimateTransformer3DModel,
HunyuanTransformer3DModel,
Transformer3DModel)
name_to_transformer3d = {
"Transformer3DModel": Transformer3DModel,
"HunyuanTransformer3DModel": HunyuanTransformer3DModel,
"EasyAnimateTransformer3DModel": EasyAnimateTransformer3DModel,
}
name_to_autoencoder_magvit = {
"AutoencoderKL": AutoencoderKL,
"AutoencoderKLMagvit": AutoencoderKLMagvit,
"AutoencoderKLCogVideoX": AutoencoderKLCogVideoX,
}
+503 -64
View File
@@ -11,25 +11,38 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import math
from typing import Any, Dict, Optional
from typing import Any, Dict, Optional, Tuple, Union
import diffusers
import pkg_resources
import torch
import torch.nn.functional as F
import torch.nn.init as init
from diffusers.models.activations import GEGLU, GELU, ApproximateGELU
from diffusers.models.attention import AdaLayerNorm, FeedForward
from diffusers.models.attention_processor import Attention
from diffusers.models.embeddings import SinusoidalPositionalEmbedding
from diffusers.models.lora import LoRACompatibleLinear
from diffusers.models.normalization import AdaLayerNorm, AdaLayerNormZero
from diffusers.utils import USE_PEFT_BACKEND
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models.attention import Attention, FeedForward
from diffusers.models.attention_processor import (Attention,
AttentionProcessor,
AttnProcessor2_0,
HunyuanAttnProcessor2_0)
from diffusers.models.embeddings import (SinusoidalPositionalEmbedding,
TimestepEmbedding, Timesteps,
get_3d_sincos_pos_embed)
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.normalization import (AdaLayerNorm, AdaLayerNormZero,
CogVideoXLayerNormZero)
from diffusers.utils import USE_PEFT_BACKEND, is_torch_version, logging
from diffusers.utils.import_utils import is_xformers_available
from diffusers.utils.torch_utils import maybe_allow_in_graph
from einops import rearrange, repeat
from torch import nn
from .motion_module import get_motion_module
from .motion_module import PositionalEncoding, get_motion_module
from .norm import AdaLayerNormShift, FP32LayerNorm, EasyAnimateLayerNormZero
from .processor import (EasyAnimateAttnProcessor2_0,
LazyKVCompressionProcessor2_0)
if is_xformers_available():
import xformers
@@ -38,6 +51,12 @@ else:
xformers = None
def zero_module(module):
# Zero out the parameters of a module and return it.
for p in module.parameters():
p.detach().zero_()
return module
@maybe_allow_in_graph
class GatedSelfAttentionDense(nn.Module):
r"""
@@ -59,8 +78,8 @@ class GatedSelfAttentionDense(nn.Module):
self.attn = Attention(query_dim=query_dim, heads=n_heads, dim_head=d_head)
self.ff = FeedForward(query_dim, activation_fn="geglu")
self.norm1 = nn.LayerNorm(query_dim)
self.norm2 = nn.LayerNorm(query_dim)
self.norm1 = FP32LayerNorm(query_dim)
self.norm2 = FP32LayerNorm(query_dim)
self.register_parameter("alpha_attn", nn.Parameter(torch.tensor(0.0)))
self.register_parameter("alpha_dense", nn.Parameter(torch.tensor(0.0)))
@@ -79,13 +98,33 @@ class GatedSelfAttentionDense(nn.Module):
return x
def zero_module(module):
# Zero out the parameters of a module and return it.
for p in module.parameters():
p.detach().zero_()
return module
class LazyKVCompressionAttention(Attention):
def __init__(
self,
sr_ratio=2, *args, **kwargs
):
super().__init__(*args, **kwargs)
self.sr_ratio = sr_ratio
self.k_compression = nn.Conv2d(
kwargs["query_dim"],
kwargs["query_dim"],
groups=kwargs["query_dim"],
kernel_size=sr_ratio,
stride=sr_ratio,
bias=True
)
self.v_compression = nn.Conv2d(
kwargs["query_dim"],
kwargs["query_dim"],
groups=kwargs["query_dim"],
kernel_size=sr_ratio,
stride=sr_ratio,
bias=True
)
init.constant_(self.k_compression.weight, 1 / (sr_ratio * sr_ratio))
init.constant_(self.v_compression.weight, 1 / (sr_ratio * sr_ratio))
init.constant_(self.k_compression.bias, 0)
init.constant_(self.v_compression.bias, 0)
@maybe_allow_in_graph
class TemporalTransformerBlock(nn.Module):
@@ -146,6 +185,8 @@ class TemporalTransformerBlock(nn.Module):
# motion module kwargs
motion_module_type = "VanillaGrid",
motion_module_kwargs = None,
qk_norm = False,
after_norm = False,
):
super().__init__()
self.only_cross_attention = only_cross_attention
@@ -178,7 +219,7 @@ class TemporalTransformerBlock(nn.Module):
elif self.use_ada_layer_norm_zero:
self.norm1 = AdaLayerNormZero(dim, num_embeds_ada_norm)
else:
self.norm1 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
self.norm1 = FP32LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
self.attn1 = Attention(
query_dim=dim,
@@ -188,6 +229,8 @@ class TemporalTransformerBlock(nn.Module):
bias=attention_bias,
cross_attention_dim=cross_attention_dim if only_cross_attention else None,
upcast_attention=upcast_attention,
qk_norm="layer_norm" if qk_norm else None,
processor=HunyuanAttnProcessor2_0() if qk_norm else AttnProcessor2_0(),
)
self.attn_temporal = get_motion_module(
@@ -204,7 +247,7 @@ class TemporalTransformerBlock(nn.Module):
self.norm2 = (
AdaLayerNorm(dim, num_embeds_ada_norm)
if self.use_ada_layer_norm
else nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
else FP32LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
)
self.attn2 = Attention(
query_dim=dim,
@@ -214,6 +257,8 @@ class TemporalTransformerBlock(nn.Module):
dropout=dropout,
bias=attention_bias,
upcast_attention=upcast_attention,
qk_norm="layer_norm" if qk_norm else None,
processor=HunyuanAttnProcessor2_0() if qk_norm else AttnProcessor2_0(),
) # is self-attn if encoder_hidden_states is none
else:
self.norm2 = None
@@ -221,10 +266,15 @@ class TemporalTransformerBlock(nn.Module):
# 3. Feed-forward
if not self.use_ada_layer_norm_single:
self.norm3 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
self.norm3 = FP32LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn, final_dropout=final_dropout)
if after_norm:
self.norm4 = FP32LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
else:
self.norm4 = None
# 4. Fuser
if attention_type == "gated" or attention_type == "gated-text-image":
self.fuser = GatedSelfAttentionDense(dim, cross_attention_dim, num_attention_heads, attention_head_dim)
@@ -330,6 +380,9 @@ class TemporalTransformerBlock(nn.Module):
if self.pos_embed is not None and self.use_ada_layer_norm_single is None:
norm_hidden_states = self.pos_embed(norm_hidden_states)
if norm_hidden_states.dtype != encoder_hidden_states.dtype or norm_hidden_states.dtype != encoder_attention_mask.dtype:
norm_hidden_states = norm_hidden_states.to(encoder_hidden_states.dtype)
attn_output = self.attn2(
norm_hidden_states,
encoder_hidden_states=encoder_hidden_states,
@@ -366,6 +419,9 @@ class TemporalTransformerBlock(nn.Module):
)
else:
ff_output = self.ff(norm_hidden_states, scale=lora_scale)
if self.norm4 is not None:
ff_output = self.norm4(ff_output)
if self.use_ada_layer_norm_zero:
ff_output = gate_mlp.unsqueeze(1) * ff_output
@@ -429,12 +485,14 @@ class SelfAttentionTemporalTransformerBlock(nn.Module):
double_self_attention: bool = False,
upcast_attention: bool = False,
norm_elementwise_affine: bool = True,
norm_type: str = "layer_norm", # 'layer_norm', 'ada_norm', 'ada_norm_zero', 'ada_norm_single'
norm_type: str = "layer_norm",
norm_eps: float = 1e-5,
final_dropout: bool = False,
attention_type: str = "default",
positional_embeddings: Optional[str] = None,
num_positional_embeddings: Optional[int] = None,
qk_norm = False,
after_norm = False,
):
super().__init__()
self.only_cross_attention = only_cross_attention
@@ -467,7 +525,7 @@ class SelfAttentionTemporalTransformerBlock(nn.Module):
elif self.use_ada_layer_norm_zero:
self.norm1 = AdaLayerNormZero(dim, num_embeds_ada_norm)
else:
self.norm1 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
self.norm1 = FP32LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
self.attn1 = Attention(
query_dim=dim,
@@ -477,6 +535,8 @@ class SelfAttentionTemporalTransformerBlock(nn.Module):
bias=attention_bias,
cross_attention_dim=cross_attention_dim if only_cross_attention else None,
upcast_attention=upcast_attention,
qk_norm="layer_norm" if qk_norm else None,
processor=HunyuanAttnProcessor2_0() if qk_norm else AttnProcessor2_0(),
)
# 2. Cross-Attn
@@ -487,7 +547,7 @@ class SelfAttentionTemporalTransformerBlock(nn.Module):
self.norm2 = (
AdaLayerNorm(dim, num_embeds_ada_norm)
if self.use_ada_layer_norm
else nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
else FP32LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
)
self.attn2 = Attention(
query_dim=dim,
@@ -497,6 +557,8 @@ class SelfAttentionTemporalTransformerBlock(nn.Module):
dropout=dropout,
bias=attention_bias,
upcast_attention=upcast_attention,
qk_norm="layer_norm" if qk_norm else None,
processor=HunyuanAttnProcessor2_0() if qk_norm else AttnProcessor2_0(),
) # is self-attn if encoder_hidden_states is none
else:
self.norm2 = None
@@ -504,10 +566,15 @@ class SelfAttentionTemporalTransformerBlock(nn.Module):
# 3. Feed-forward
if not self.use_ada_layer_norm_single:
self.norm3 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
self.norm3 = FP32LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn, final_dropout=final_dropout)
if after_norm:
self.norm4 = FP32LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
else:
self.norm4 = None
# 4. Fuser
if attention_type == "gated" or attention_type == "gated-text-image":
self.fuser = GatedSelfAttentionDense(dim, cross_attention_dim, num_attention_heads, attention_head_dim)
@@ -639,6 +706,9 @@ class SelfAttentionTemporalTransformerBlock(nn.Module):
)
else:
ff_output = self.ff(norm_hidden_states, scale=lora_scale)
if self.norm4 is not None:
ff_output = self.norm4(ff_output)
if self.use_ada_layer_norm_zero:
ff_output = gate_mlp.unsqueeze(1) * ff_output
@@ -651,59 +721,428 @@ class SelfAttentionTemporalTransformerBlock(nn.Module):
return hidden_states
class GEGLU(nn.Module):
def __init__(self, dim_in, dim_out, norm_elementwise_affine):
super().__init__()
self.norm = FP32LayerNorm(dim_in, dim_in, norm_elementwise_affine)
self.proj = nn.Linear(dim_in, dim_out * 2)
class FeedForward(nn.Module):
def forward(self, x):
x, gate = self.proj(self.norm(x)).chunk(2, dim=-1)
return x * F.gelu(gate)
@maybe_allow_in_graph
class HunyuanDiTBlock(nn.Module):
r"""
A feed-forward layer.
Transformer block used in Hunyuan-DiT model (https://github.com/Tencent/HunyuanDiT). Allow skip connection and
QKNorm
Parameters:
dim (`int`): The number of channels in the input.
dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`.
mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
final_dropout (`bool` *optional*, defaults to False): Apply a final dropout.
dim (`int`):
The number of channels in the input and output.
num_attention_heads (`int`):
The number of headsto use for multi-head attention.
cross_attention_dim (`int`,*optional*):
The size of the encoder_hidden_states vector for cross attention.
dropout(`float`, *optional*, defaults to 0.0):
The dropout probability to use.
activation_fn (`str`,*optional*, defaults to `"geglu"`):
Activation function to be used in feed-forward. .
norm_elementwise_affine (`bool`, *optional*, defaults to `True`):
Whether to use learnable elementwise affine parameters for normalization.
norm_eps (`float`, *optional*, defaults to 1e-6):
A small constant added to the denominator in normalization layers to prevent division by zero.
final_dropout (`bool` *optional*, defaults to False):
Whether to apply a final dropout after the last feed-forward layer.
ff_inner_dim (`int`, *optional*):
The size of the hidden layer in the feed-forward block. Defaults to `None`.
ff_bias (`bool`, *optional*, defaults to `True`):
Whether to use bias in the feed-forward block.
skip (`bool`, *optional*, defaults to `False`):
Whether to use skip connection. Defaults to `False` for down-blocks and mid-blocks.
qk_norm (`bool`, *optional*, defaults to `True`):
Whether to use normalization in QK calculation. Defaults to `True`.
"""
def __init__(
self,
dim: int,
dim_out: Optional[int] = None,
mult: int = 4,
dropout: float = 0.0,
num_attention_heads: int,
cross_attention_dim: int = 1024,
dropout=0.0,
activation_fn: str = "geglu",
norm_elementwise_affine: bool = True,
norm_eps: float = 1e-6,
final_dropout: bool = False,
ff_inner_dim: Optional[int] = None,
ff_bias: bool = True,
skip: bool = False,
qk_norm: bool = True,
time_position_encoding: bool = False,
after_norm: bool = False,
is_local_attention: bool = False,
local_attention_frames: int = 2,
enable_inpaint: bool = False,
kvcompression = False,
):
super().__init__()
inner_dim = int(dim * mult)
dim_out = dim_out if dim_out is not None else dim
linear_cls = LoRACompatibleLinear if not USE_PEFT_BACKEND else nn.Linear
if activation_fn == "gelu":
act_fn = GELU(dim, inner_dim)
if activation_fn == "gelu-approximate":
act_fn = GELU(dim, inner_dim, approximate="tanh")
elif activation_fn == "geglu":
act_fn = GEGLU(dim, inner_dim)
elif activation_fn == "geglu-approximate":
act_fn = ApproximateGELU(dim, inner_dim)
# Define 3 blocks. Each block has its own normalization layer.
# NOTE: when new version comes, check norm2 and norm 3
# 1. Self-Attn
self.norm1 = AdaLayerNormShift(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
self.t_embed = PositionalEncoding(dim, dropout=0., max_len=512) \
if time_position_encoding else nn.Identity()
self.net = nn.ModuleList([])
# project in
self.net.append(act_fn)
# project dropout
self.net.append(nn.Dropout(dropout))
# project out
self.net.append(linear_cls(inner_dim, dim_out))
# FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout
if final_dropout:
self.net.append(nn.Dropout(dropout))
self.is_local_attention = is_local_attention
self.local_attention_frames = local_attention_frames
self.kvcompression = kvcompression
if kvcompression:
self.attn1 = LazyKVCompressionAttention(
query_dim=dim,
cross_attention_dim=None,
dim_head=dim // num_attention_heads,
heads=num_attention_heads,
qk_norm="layer_norm" if qk_norm else None,
eps=1e-6,
bias=True,
processor=LazyKVCompressionProcessor2_0(),
)
else:
self.attn1 = Attention(
query_dim=dim,
cross_attention_dim=None,
dim_head=dim // num_attention_heads,
heads=num_attention_heads,
qk_norm="layer_norm" if qk_norm else None,
eps=1e-6,
bias=True,
processor=HunyuanAttnProcessor2_0(),
)
def forward(self, hidden_states: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
compatible_cls = (GEGLU,) if USE_PEFT_BACKEND else (GEGLU, LoRACompatibleLinear)
for module in self.net:
if isinstance(module, compatible_cls):
hidden_states = module(hidden_states, scale)
# 2. Cross-Attn
self.norm2 = FP32LayerNorm(dim, norm_eps, norm_elementwise_affine)
if self.is_local_attention:
from mamba_ssm import Mamba2
self.mamba_norm_in = FP32LayerNorm(dim, norm_eps, norm_elementwise_affine)
self.in_linear = nn.Linear(dim, 1536)
self.mamba_norm_1 = FP32LayerNorm(1536, norm_eps, norm_elementwise_affine)
self.mamba_norm_2 = FP32LayerNorm(1536, norm_eps, norm_elementwise_affine)
self.mamba_block_1 = Mamba2(
d_model=1536,
d_state=64,
d_conv=4,
expand=2,
)
self.mamba_block_2 = Mamba2(
d_model=1536,
d_state=64,
d_conv=4,
expand=2,
)
self.mamba_norm_after_mamba_block = FP32LayerNorm(1536, norm_eps, norm_elementwise_affine)
self.out_linear = nn.Linear(1536, dim)
self.out_linear = zero_module(self.out_linear)
self.mamba_norm_out = FP32LayerNorm(dim, norm_eps, norm_elementwise_affine)
self.attn2 = Attention(
query_dim=dim,
cross_attention_dim=cross_attention_dim,
dim_head=dim // num_attention_heads,
heads=num_attention_heads,
qk_norm="layer_norm" if qk_norm else None,
eps=1e-6,
bias=True,
processor=HunyuanAttnProcessor2_0(),
)
if enable_inpaint:
self.norm_clip = FP32LayerNorm(dim, norm_eps, norm_elementwise_affine)
self.attn_clip = Attention(
query_dim=dim,
cross_attention_dim=cross_attention_dim,
dim_head=dim // num_attention_heads,
heads=num_attention_heads,
qk_norm="layer_norm" if qk_norm else None,
eps=1e-6,
bias=True,
processor=HunyuanAttnProcessor2_0(),
)
self.gate_clip = GEGLU(dim, dim, norm_elementwise_affine)
self.norm_clip_out = FP32LayerNorm(dim, norm_eps, norm_elementwise_affine)
else:
self.attn_clip = None
self.norm_clip = None
self.gate_clip = None
self.norm_clip_out = None
# 3. Feed-forward
self.norm3 = FP32LayerNorm(dim, norm_eps, norm_elementwise_affine)
self.ff = FeedForward(
dim,
dropout=dropout, ### 0.0
activation_fn=activation_fn, ### approx GeLU
final_dropout=final_dropout, ### 0.0
inner_dim=ff_inner_dim, ### int(dim * mlp_ratio)
bias=ff_bias,
)
# 4. Skip Connection
if skip:
self.skip_norm = FP32LayerNorm(2 * dim, norm_eps, elementwise_affine=True)
self.skip_linear = nn.Linear(2 * dim, dim)
else:
self.skip_linear = None
if after_norm:
self.norm4 = FP32LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
else:
self.norm4 = None
# let chunk size default to None
self._chunk_size = None
self._chunk_dim = 0
def set_chunk_feed_forward(self, chunk_size: Optional[int], dim: int = 0):
# Sets chunk feed-forward
self._chunk_size = chunk_size
self._chunk_dim = dim
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
temb: Optional[torch.Tensor] = None,
image_rotary_emb=None,
skip=None,
num_frames: int = 1,
height: int = 32,
width: int = 32,
clip_encoder_hidden_states: Optional[torch.Tensor] = None,
disable_image_rotary_emb_in_attn1=False,
) -> torch.Tensor:
# Notice that normalization is always applied before the real computation in the following blocks.
# 0. Long Skip Connection
if self.skip_linear is not None:
cat = torch.cat([hidden_states, skip], dim=-1)
cat = self.skip_norm(cat)
hidden_states = self.skip_linear(cat)
if image_rotary_emb is not None:
image_rotary_emb = (torch.cat([image_rotary_emb[0] for i in range(num_frames)], dim=0), torch.cat([image_rotary_emb[1] for i in range(num_frames)], dim=0))
if num_frames != 1:
# add time embedding
hidden_states = rearrange(hidden_states, "b (f d) c -> (b d) f c", f=num_frames)
if self.t_embed is not None:
hidden_states = self.t_embed(hidden_states)
hidden_states = rearrange(hidden_states, "(b d) f c -> b (f d) c", d=height * width)
# 1. Self-Attention
norm_hidden_states = self.norm1(hidden_states, temb) ### checked: self.norm1 is correct
if num_frames > 2 and self.is_local_attention:
if image_rotary_emb is not None:
attn1_image_rotary_emb = (image_rotary_emb[0][:int(height * width * 2)], image_rotary_emb[1][:int(height * width * 2)])
else:
hidden_states = module(hidden_states)
attn1_image_rotary_emb = image_rotary_emb
norm_hidden_states_1 = rearrange(norm_hidden_states, "b (f d) c -> b f d c", d=height * width)
norm_hidden_states_1 = rearrange(norm_hidden_states_1, "b (f p) d c -> (b f) (p d) c", p = 2)
attn_output = self.attn1(
norm_hidden_states_1,
image_rotary_emb=attn1_image_rotary_emb if not disable_image_rotary_emb_in_attn1 else None,
)
attn_output = rearrange(attn_output, "(b f) (p d) c -> b (f p) d c", p = 2, f = num_frames // 2)
norm_hidden_states_2 = rearrange(norm_hidden_states, "b (f d) c -> b f d c", d = height * width)[:, 1:-1]
local_attention_frames_num = norm_hidden_states_2.size()[1] // 2
norm_hidden_states_2 = rearrange(norm_hidden_states_2, "b (f p) d c -> (b f) (p d) c", p = 2)
attn_output_2 = self.attn1(
norm_hidden_states_2,
image_rotary_emb=attn1_image_rotary_emb if not disable_image_rotary_emb_in_attn1 else None,
)
attn_output_2 = rearrange(attn_output_2, "(b f) (p d) c -> b (f p) d c", p = 2, f = local_attention_frames_num)
attn_output[:, 1:-1] = (attn_output[:, 1:-1] + attn_output_2) / 2
attn_output = rearrange(attn_output, "b f d c -> b (f d) c")
else:
if self.kvcompression:
norm_hidden_states = rearrange(norm_hidden_states, "b (f h w) c -> b c f h w", f = num_frames, h = height, w = width)
attn_output = self.attn1(
norm_hidden_states,
image_rotary_emb=image_rotary_emb if not disable_image_rotary_emb_in_attn1 else None,
)
else:
attn_output = self.attn1(
norm_hidden_states,
image_rotary_emb=image_rotary_emb if not disable_image_rotary_emb_in_attn1 else None,
)
hidden_states = hidden_states + attn_output
if num_frames > 2 and self.is_local_attention:
hidden_states_in = self.in_linear(self.mamba_norm_in(hidden_states))
hidden_states = hidden_states + self.mamba_norm_out(
self.out_linear(
self.mamba_norm_after_mamba_block(
self.mamba_block_1(
self.mamba_norm_1(hidden_states_in)
) +
self.mamba_block_2(
self.mamba_norm_2(hidden_states_in.flip(1))
).flip(1)
)
)
)
# 2. Cross-Attention
hidden_states = hidden_states + self.attn2(
self.norm2(hidden_states),
encoder_hidden_states=encoder_hidden_states,
image_rotary_emb=image_rotary_emb,
)
if self.attn_clip is not None:
hidden_states = hidden_states + self.norm_clip_out(
self.gate_clip(
self.attn_clip(
self.norm_clip(hidden_states),
encoder_hidden_states=clip_encoder_hidden_states,
image_rotary_emb=image_rotary_emb,
)
)
)
# FFN Layer ### TODO: switch norm2 and norm3 in the state dict
mlp_inputs = self.norm3(hidden_states)
if self.norm4 is not None:
hidden_states = hidden_states + self.norm4(self.ff(mlp_inputs))
else:
hidden_states = hidden_states + self.ff(mlp_inputs)
return hidden_states
@maybe_allow_in_graph
class EasyAnimateDiTBlock(nn.Module):
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
time_embed_dim: int,
dropout: float = 0.0,
activation_fn: str = "gelu-approximate",
norm_elementwise_affine: bool = True,
norm_eps: float = 1e-6,
final_dropout: bool = True,
ff_inner_dim: Optional[int] = None,
ff_bias: bool = True,
qk_norm: bool = True,
after_norm: bool = False,
norm_type: str="fp32_layer_norm",
is_mmdit_block: bool = True,
):
super().__init__()
# Attention Part
self.norm1 = EasyAnimateLayerNormZero(
time_embed_dim, dim, norm_elementwise_affine, norm_eps, norm_type=norm_type, bias=True
)
self.attn1 = Attention(
query_dim=dim,
dim_head=attention_head_dim,
heads=num_attention_heads,
qk_norm="layer_norm" if qk_norm else None,
eps=1e-6,
bias=True,
processor=EasyAnimateAttnProcessor2_0(),
)
if is_mmdit_block:
self.attn2 = Attention(
query_dim=dim,
dim_head=attention_head_dim,
heads=num_attention_heads,
qk_norm="layer_norm" if qk_norm else None,
eps=1e-6,
bias=True,
processor=EasyAnimateAttnProcessor2_0(),
)
else:
self.attn2 = None
# FFN Part
self.norm2 = EasyAnimateLayerNormZero(
time_embed_dim, dim, norm_elementwise_affine, norm_eps, norm_type=norm_type, bias=True
)
self.ff = FeedForward(
dim,
dropout=dropout,
activation_fn=activation_fn,
final_dropout=final_dropout,
inner_dim=ff_inner_dim,
bias=ff_bias,
)
if is_mmdit_block:
self.txt_ff = FeedForward(
dim,
dropout=dropout,
activation_fn=activation_fn,
final_dropout=final_dropout,
inner_dim=ff_inner_dim,
bias=ff_bias,
)
else:
self.txt_ff = None
if after_norm:
self.norm3 = FP32LayerNorm(dim, elementwise_affine=norm_elementwise_affine, eps=norm_eps)
else:
self.norm3 = None
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> torch.Tensor:
# Norm
norm_hidden_states, norm_encoder_hidden_states, gate_msa, enc_gate_msa = self.norm1(
hidden_states, encoder_hidden_states, temb
)
# Attn
attn_hidden_states, attn_encoder_hidden_states = self.attn1(
hidden_states=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
image_rotary_emb=image_rotary_emb,
attn2=self.attn2,
)
hidden_states = hidden_states + gate_msa * attn_hidden_states
encoder_hidden_states = encoder_hidden_states + enc_gate_msa * attn_encoder_hidden_states
# Norm
norm_hidden_states, norm_encoder_hidden_states, gate_ff, enc_gate_ff = self.norm2(
hidden_states, encoder_hidden_states, temb
)
# FFN
if self.norm3 is not None:
norm_hidden_states = self.norm3(self.ff(norm_hidden_states))
if self.txt_ff is not None:
norm_encoder_hidden_states = self.norm3(self.txt_ff(norm_encoder_hidden_states))
else:
norm_encoder_hidden_states = self.norm3(self.ff(norm_encoder_hidden_states))
else:
norm_hidden_states = self.ff(norm_hidden_states)
if self.txt_ff is not None:
norm_encoder_hidden_states = self.txt_ff(norm_encoder_hidden_states)
else:
norm_encoder_hidden_states = self.ff(norm_encoder_hidden_states)
hidden_states = hidden_states + gate_ff * norm_hidden_states
encoder_hidden_states = encoder_hidden_states + enc_gate_ff * norm_encoder_hidden_states
return hidden_states, encoder_hidden_states
File diff suppressed because it is too large Load Diff
+107
View File
@@ -0,0 +1,107 @@
import math
from typing import Optional
import numpy as np
import torch
import torch.nn.functional as F
from diffusers.models.embeddings import (PixArtAlphaTextProjection, get_timestep_embedding,
TimestepEmbedding, Timesteps)
from einops import rearrange
from torch import nn
class HunyuanDiTAttentionPool(nn.Module):
def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: int = None):
super().__init__()
self.positional_embedding = nn.Parameter(torch.randn(spacial_dim + 1, embed_dim) / embed_dim**0.5)
self.k_proj = nn.Linear(embed_dim, embed_dim)
self.q_proj = nn.Linear(embed_dim, embed_dim)
self.v_proj = nn.Linear(embed_dim, embed_dim)
self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim)
self.num_heads = num_heads
def forward(self, x):
x = torch.cat([x.mean(dim=1, keepdim=True), x], dim=1)
x = x + self.positional_embedding[None, :, :].to(x.dtype)
query = self.q_proj(x[:, :1])
key = self.k_proj(x)
value = self.v_proj(x)
batch_size, _, _ = query.size()
query = query.reshape(batch_size, -1, self.num_heads, query.size(-1) // self.num_heads).transpose(1, 2) # (1, H, N, E/H)
key = key.reshape(batch_size, -1, self.num_heads, key.size(-1) // self.num_heads).transpose(1, 2) # (L+1, H, N, E/H)
value = value.reshape(batch_size, -1, self.num_heads, value.size(-1) // self.num_heads).transpose(1, 2) # (L+1, H, N, E/H)
x = F.scaled_dot_product_attention(query=query, key=key, value=value, attn_mask=None, dropout_p=0.0, is_causal=False)
x = x.transpose(1, 2).reshape(batch_size, 1, -1)
x = x.to(query.dtype)
x = self.c_proj(x)
return x.squeeze(1)
class HunyuanCombinedTimestepTextSizeStyleEmbedding(nn.Module):
def __init__(self, embedding_dim, pooled_projection_dim=1024, seq_len=256, cross_attention_dim=2048):
super().__init__()
self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)
self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)
self.pooler = HunyuanDiTAttentionPool(
seq_len, cross_attention_dim, num_heads=8, output_dim=pooled_projection_dim
)
# Here we use a default learned embedder layer for future extension.
self.style_embedder = nn.Embedding(1, embedding_dim)
extra_in_dim = 256 * 6 + embedding_dim + pooled_projection_dim
self.extra_embedder = PixArtAlphaTextProjection(
in_features=extra_in_dim,
hidden_size=embedding_dim * 4,
out_features=embedding_dim,
act_fn="silu_fp32",
)
def forward(self, timestep, encoder_hidden_states, image_meta_size, style, hidden_dtype=None):
timesteps_proj = self.time_proj(timestep)
timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, 256)
# extra condition1: text
pooled_projections = self.pooler(encoder_hidden_states) # (N, 1024)
# extra condition2: image meta size embdding
image_meta_size = get_timestep_embedding(image_meta_size.view(-1), 256, True, 0)
image_meta_size = image_meta_size.to(dtype=hidden_dtype)
image_meta_size = image_meta_size.view(-1, 6 * 256) # (N, 1536)
# extra condition3: style embedding
style_embedding = self.style_embedder(style) # (N, embedding_dim)
# Concatenate all extra vectors
extra_cond = torch.cat([pooled_projections, image_meta_size, style_embedding], dim=1)
conditioning = timesteps_emb + self.extra_embedder(extra_cond) # [B, D]
return conditioning
class TimePositionalEncoding(nn.Module):
def __init__(
self,
d_model,
dropout = 0.,
max_len = 24
):
super().__init__()
self.dropout = nn.Dropout(p=dropout)
position = torch.arange(max_len).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
pe = torch.zeros(1, max_len, d_model)
pe[0, :, 0::2] = torch.sin(position * div_term)
pe[0, :, 1::2] = torch.cos(position * div_term)
self.register_buffer('pe', pe)
def forward(self, x):
b, c, f, h, w = x.size()
x = rearrange(x, "b c f h w -> (b h w) f c")
x = x + self.pe[:, :x.size(1)]
x = rearrange(x, "(b h w) f c -> b c f h w", b=b, h=h, w=w)
return self.dropout(x)
+146 -277
View File
@@ -1,248 +1,33 @@
"""Modified from https://github.com/guoyww/AnimateDiff/blob/main/animatediff/models/motion_module.py
"""
import math
from typing import Any, Callable, List, Optional, Tuple, Union
import diffusers
import pkg_resources
import torch
import torch.nn.functional as F
installed_version = diffusers.__version__
if pkg_resources.parse_version(installed_version) >= pkg_resources.parse_version("0.28.2"):
from diffusers.models.attention_processor import (Attention,
AttnProcessor2_0,
HunyuanAttnProcessor2_0)
else:
from diffusers.models.attention_processor import Attention, AttnProcessor2_0
from diffusers.models.attention import FeedForward
from diffusers.utils.import_utils import is_xformers_available
from einops import rearrange, repeat
from torch import nn
from .norm import FP32LayerNorm
if is_xformers_available():
import xformers
import xformers.ops
else:
xformers = None
class CrossAttention(nn.Module):
r"""
A cross attention layer.
Parameters:
query_dim (`int`): The number of channels in the query.
cross_attention_dim (`int`, *optional*):
The number of channels in the encoder_hidden_states. If not given, defaults to `query_dim`.
heads (`int`, *optional*, defaults to 8): The number of heads to use for multi-head attention.
dim_head (`int`, *optional*, defaults to 64): The number of channels in each head.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
bias (`bool`, *optional*, defaults to False):
Set to `True` for the query, key, and value linear layers to contain a bias parameter.
"""
def __init__(
self,
query_dim: int,
cross_attention_dim: Optional[int] = None,
heads: int = 8,
dim_head: int = 64,
dropout: float = 0.0,
bias=False,
upcast_attention: bool = False,
upcast_softmax: bool = False,
added_kv_proj_dim: Optional[int] = None,
norm_num_groups: Optional[int] = None,
):
super().__init__()
inner_dim = dim_head * heads
cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim
self.upcast_attention = upcast_attention
self.upcast_softmax = upcast_softmax
self.scale = dim_head**-0.5
self.heads = heads
# for slice_size > 0 the attention score computation
# is split across the batch axis to save memory
# You can set slice_size with `set_attention_slice`
self.sliceable_head_dim = heads
self._slice_size = None
self._use_memory_efficient_attention_xformers = False
self.added_kv_proj_dim = added_kv_proj_dim
if norm_num_groups is not None:
self.group_norm = nn.GroupNorm(num_channels=inner_dim, num_groups=norm_num_groups, eps=1e-5, affine=True)
else:
self.group_norm = None
self.to_q = nn.Linear(query_dim, inner_dim, bias=bias)
self.to_k = nn.Linear(cross_attention_dim, inner_dim, bias=bias)
self.to_v = nn.Linear(cross_attention_dim, inner_dim, bias=bias)
if self.added_kv_proj_dim is not None:
self.add_k_proj = nn.Linear(added_kv_proj_dim, cross_attention_dim)
self.add_v_proj = nn.Linear(added_kv_proj_dim, cross_attention_dim)
self.to_out = nn.ModuleList([])
self.to_out.append(nn.Linear(inner_dim, query_dim))
self.to_out.append(nn.Dropout(dropout))
def set_use_memory_efficient_attention_xformers(
self, valid: bool, attention_op: Optional[Callable] = None
) -> None:
self._use_memory_efficient_attention_xformers = valid
def reshape_heads_to_batch_dim(self, tensor):
batch_size, seq_len, dim = tensor.shape
head_size = self.heads
tensor = tensor.reshape(batch_size, seq_len, head_size, dim // head_size)
tensor = tensor.permute(0, 2, 1, 3).reshape(batch_size * head_size, seq_len, dim // head_size)
return tensor
def reshape_batch_dim_to_heads(self, tensor):
batch_size, seq_len, dim = tensor.shape
head_size = self.heads
tensor = tensor.reshape(batch_size // head_size, head_size, seq_len, dim)
tensor = tensor.permute(0, 2, 1, 3).reshape(batch_size // head_size, seq_len, dim * head_size)
return tensor
def set_attention_slice(self, slice_size):
if slice_size is not None and slice_size > self.sliceable_head_dim:
raise ValueError(f"slice_size {slice_size} has to be smaller or equal to {self.sliceable_head_dim}.")
self._slice_size = slice_size
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None):
batch_size, sequence_length, _ = hidden_states.shape
encoder_hidden_states = encoder_hidden_states
if self.group_norm is not None:
hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = self.to_q(hidden_states)
dim = query.shape[-1]
query = self.reshape_heads_to_batch_dim(query)
if self.added_kv_proj_dim is not None:
key = self.to_k(hidden_states)
value = self.to_v(hidden_states)
encoder_hidden_states_key_proj = self.add_k_proj(encoder_hidden_states)
encoder_hidden_states_value_proj = self.add_v_proj(encoder_hidden_states)
key = self.reshape_heads_to_batch_dim(key)
value = self.reshape_heads_to_batch_dim(value)
encoder_hidden_states_key_proj = self.reshape_heads_to_batch_dim(encoder_hidden_states_key_proj)
encoder_hidden_states_value_proj = self.reshape_heads_to_batch_dim(encoder_hidden_states_value_proj)
key = torch.concat([encoder_hidden_states_key_proj, key], dim=1)
value = torch.concat([encoder_hidden_states_value_proj, value], dim=1)
else:
encoder_hidden_states = encoder_hidden_states if encoder_hidden_states is not None else hidden_states
key = self.to_k(encoder_hidden_states)
value = self.to_v(encoder_hidden_states)
key = self.reshape_heads_to_batch_dim(key)
value = self.reshape_heads_to_batch_dim(value)
if attention_mask is not None:
if attention_mask.shape[-1] != query.shape[1]:
target_length = query.shape[1]
attention_mask = F.pad(attention_mask, (0, target_length), value=0.0)
attention_mask = attention_mask.repeat_interleave(self.heads, dim=0)
# attention, what we cannot get enough of
if self._use_memory_efficient_attention_xformers:
hidden_states = self._memory_efficient_attention_xformers(query, key, value, attention_mask)
# Some versions of xformers return output in fp32, cast it back to the dtype of the input
hidden_states = hidden_states.to(query.dtype)
else:
if self._slice_size is None or query.shape[0] // self._slice_size == 1:
hidden_states = self._attention(query, key, value, attention_mask)
else:
hidden_states = self._sliced_attention(query, key, value, sequence_length, dim, attention_mask)
# linear proj
hidden_states = self.to_out[0](hidden_states)
# dropout
hidden_states = self.to_out[1](hidden_states)
return hidden_states
def _attention(self, query, key, value, attention_mask=None):
if self.upcast_attention:
query = query.float()
key = key.float()
attention_scores = torch.baddbmm(
torch.empty(query.shape[0], query.shape[1], key.shape[1], dtype=query.dtype, device=query.device),
query,
key.transpose(-1, -2),
beta=0,
alpha=self.scale,
)
if attention_mask is not None:
attention_scores = attention_scores + attention_mask
if self.upcast_softmax:
attention_scores = attention_scores.float()
attention_probs = attention_scores.softmax(dim=-1)
# cast back to the original dtype
attention_probs = attention_probs.to(value.dtype)
# compute attention output
hidden_states = torch.bmm(attention_probs, value)
# reshape hidden_states
hidden_states = self.reshape_batch_dim_to_heads(hidden_states)
return hidden_states
def _sliced_attention(self, query, key, value, sequence_length, dim, attention_mask):
batch_size_attention = query.shape[0]
hidden_states = torch.zeros(
(batch_size_attention, sequence_length, dim // self.heads), device=query.device, dtype=query.dtype
)
slice_size = self._slice_size if self._slice_size is not None else hidden_states.shape[0]
for i in range(hidden_states.shape[0] // slice_size):
start_idx = i * slice_size
end_idx = (i + 1) * slice_size
query_slice = query[start_idx:end_idx]
key_slice = key[start_idx:end_idx]
if self.upcast_attention:
query_slice = query_slice.float()
key_slice = key_slice.float()
attn_slice = torch.baddbmm(
torch.empty(slice_size, query.shape[1], key.shape[1], dtype=query_slice.dtype, device=query.device),
query_slice,
key_slice.transpose(-1, -2),
beta=0,
alpha=self.scale,
)
if attention_mask is not None:
attn_slice = attn_slice + attention_mask[start_idx:end_idx]
if self.upcast_softmax:
attn_slice = attn_slice.float()
attn_slice = attn_slice.softmax(dim=-1)
# cast back to the original dtype
attn_slice = attn_slice.to(value.dtype)
attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx])
hidden_states[start_idx:end_idx] = attn_slice
# reshape hidden_states
hidden_states = self.reshape_batch_dim_to_heads(hidden_states)
return hidden_states
def _memory_efficient_attention_xformers(self, query, key, value, attention_mask):
# TODO attention_mask
query = query.contiguous()
key = key.contiguous()
value = value.contiguous()
hidden_states = xformers.ops.memory_efficient_attention(query, key, value, attn_bias=attention_mask)
hidden_states = self.reshape_batch_dim_to_heads(hidden_states)
return hidden_states
def zero_module(module):
# Zero out the parameters of a module and return it.
for p in module.parameters():
@@ -275,6 +60,11 @@ class VanillaTemporalModule(nn.Module):
zero_initialize = True,
block_size = 1,
grid = False,
remove_time_embedding_in_photo = False,
global_num_attention_heads = 16,
global_attention = False,
qk_norm = False,
):
super().__init__()
@@ -289,17 +79,87 @@ class VanillaTemporalModule(nn.Module):
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
grid=grid,
block_size=block_size,
remove_time_embedding_in_photo=remove_time_embedding_in_photo,
qk_norm=qk_norm,
)
self.global_transformer = GlobalTransformer3DModel(
in_channels=in_channels,
num_attention_heads=global_num_attention_heads,
attention_head_dim=in_channels // global_num_attention_heads // temporal_attention_dim_div,
qk_norm=qk_norm,
) if global_attention else None
if zero_initialize:
self.temporal_transformer.proj_out = zero_module(self.temporal_transformer.proj_out)
if global_attention:
self.global_transformer.proj_out = zero_module(self.global_transformer.proj_out)
def forward(self, input_tensor, encoder_hidden_states=None, attention_mask=None, anchor_frame_idx=None):
hidden_states = input_tensor
hidden_states = self.temporal_transformer(hidden_states, encoder_hidden_states, attention_mask)
if self.global_transformer is not None:
hidden_states = self.global_transformer(hidden_states)
output = hidden_states
return output
class GlobalTransformer3DModel(nn.Module):
def __init__(
self,
in_channels,
num_attention_heads,
attention_head_dim,
dropout = 0.0,
attention_bias = False,
upcast_attention = False,
qk_norm = False,
):
super().__init__()
inner_dim = num_attention_heads * attention_head_dim
self.norm1 = FP32LayerNorm(inner_dim)
self.proj_in = nn.Linear(in_channels, inner_dim)
self.norm2 = FP32LayerNorm(inner_dim)
if pkg_resources.parse_version(installed_version) >= pkg_resources.parse_version("0.28.2"):
self.attention = Attention(
query_dim=inner_dim,
heads=num_attention_heads,
dim_head=attention_head_dim,
dropout=dropout,
bias=attention_bias,
upcast_attention=upcast_attention,
qk_norm="layer_norm" if qk_norm else None,
processor=HunyuanAttnProcessor2_0() if qk_norm else AttnProcessor2_0(),
)
else:
self.attention = Attention(
query_dim=inner_dim,
heads=num_attention_heads,
dim_head=attention_head_dim,
dropout=dropout,
bias=attention_bias,
upcast_attention=upcast_attention,
)
self.proj_out = nn.Linear(inner_dim, in_channels)
def forward(self, hidden_states):
assert hidden_states.dim() == 5, f"Expected hidden_states to have ndim=5, but got ndim={hidden_states.dim()}."
video_length, height, width = hidden_states.shape[2], hidden_states.shape[3], hidden_states.shape[4]
hidden_states = rearrange(hidden_states, "b c f h w -> b (f h w) c")
residual = hidden_states
hidden_states = self.norm1(hidden_states)
hidden_states = self.proj_in(hidden_states)
# Attention Blocks
hidden_states = self.norm2(hidden_states)
hidden_states = self.attention(hidden_states)
hidden_states = self.proj_out(hidden_states)
output = hidden_states + residual
output = rearrange(output, "b (f h w) c -> b c f h w", f=video_length, h=height, w=width)
return output
class TemporalTransformer3DModel(nn.Module):
def __init__(
self,
@@ -321,6 +181,8 @@ class TemporalTransformer3DModel(nn.Module):
temporal_position_encoding_max_len = 4096,
grid = False,
block_size = 1,
remove_time_embedding_in_photo = False,
qk_norm = False,
):
super().__init__()
@@ -348,6 +210,8 @@ class TemporalTransformer3DModel(nn.Module):
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
block_size=block_size,
grid=grid,
remove_time_embedding_in_photo=remove_time_embedding_in_photo,
qk_norm=qk_norm
)
for d in range(num_layers)
]
@@ -398,6 +262,8 @@ class TemporalTransformerBlock(nn.Module):
temporal_position_encoding_max_len = 4096,
block_size = 1,
grid = False,
remove_time_embedding_in_photo = False,
qk_norm = False,
):
super().__init__()
@@ -422,15 +288,36 @@ class TemporalTransformerBlock(nn.Module):
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
block_size=block_size,
grid=grid,
remove_time_embedding_in_photo=remove_time_embedding_in_photo,
qk_norm="layer_norm" if qk_norm else None,
processor=HunyuanAttnProcessor2_0() if qk_norm else AttnProcessor2_0(),
) if pkg_resources.parse_version(installed_version) >= pkg_resources.parse_version("0.28.2") else \
VersatileAttention(
attention_mode=block_name.split("_")[0],
cross_attention_dim=cross_attention_dim if block_name.endswith("_Cross") else None,
query_dim=dim,
heads=num_attention_heads,
dim_head=attention_head_dim,
dropout=dropout,
bias=attention_bias,
upcast_attention=upcast_attention,
cross_frame_attention_mode=cross_frame_attention_mode,
temporal_position_encoding=temporal_position_encoding,
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
block_size=block_size,
grid=grid,
remove_time_embedding_in_photo=remove_time_embedding_in_photo,
)
)
norms.append(nn.LayerNorm(dim))
norms.append(FP32LayerNorm(dim))
self.attention_blocks = nn.ModuleList(attention_blocks)
self.norms = nn.ModuleList(norms)
self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn)
self.ff_norm = nn.LayerNorm(dim)
self.ff_norm = FP32LayerNorm(dim)
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, video_length=None, height=None, weight=None):
for attention_block, norm in zip(self.attention_blocks, self.norms):
@@ -468,7 +355,7 @@ class PositionalEncoding(nn.Module):
x = x + self.pe[:, :x.size(1)]
return self.dropout(x)
class VersatileAttention(CrossAttention):
class VersatileAttention(Attention):
def __init__(
self,
attention_mode = None,
@@ -477,21 +364,23 @@ class VersatileAttention(CrossAttention):
temporal_position_encoding_max_len = 4096,
grid = False,
block_size = 1,
remove_time_embedding_in_photo = False,
*args, **kwargs
):
super().__init__(*args, **kwargs)
assert attention_mode == "Temporal"
assert attention_mode == "Temporal" or attention_mode == "Global"
self.attention_mode = attention_mode
self.is_cross_attention = kwargs["cross_attention_dim"] is not None
self.block_size = block_size
self.grid = grid
self.remove_time_embedding_in_photo = remove_time_embedding_in_photo
self.pos_encoder = PositionalEncoding(
kwargs["query_dim"],
dropout=0.,
max_len=temporal_position_encoding_max_len
) if (temporal_position_encoding and attention_mode == "Temporal") else None
) if (temporal_position_encoding and attention_mode == "Temporal") or (temporal_position_encoding and attention_mode == "Global") else None
def extra_repr(self):
return f"(Module Info) Attention_Mode: {self.attention_mode}, Is_Cross_Attention: {self.is_cross_attention}"
@@ -503,8 +392,13 @@ class VersatileAttention(CrossAttention):
# for add pos_encoder
_, before_d, _c = hidden_states.size()
hidden_states = rearrange(hidden_states, "(b f) d c -> (b d) f c", f=video_length)
if self.pos_encoder is not None:
hidden_states = self.pos_encoder(hidden_states)
if self.remove_time_embedding_in_photo:
if self.pos_encoder is not None and video_length > 1:
hidden_states = self.pos_encoder(hidden_states)
else:
if self.pos_encoder is not None:
hidden_states = self.pos_encoder(hidden_states)
if self.grid:
hidden_states = rearrange(hidden_states, "(b d) f c -> b f d c", f=video_length, d=before_d)
@@ -515,61 +409,36 @@ class VersatileAttention(CrossAttention):
else:
d = before_d
encoder_hidden_states = repeat(encoder_hidden_states, "b n c -> (b d) n c", d=d) if encoder_hidden_states is not None else encoder_hidden_states
elif self.attention_mode == "Global":
# for add pos_encoder
_, d, _c = hidden_states.size()
hidden_states = rearrange(hidden_states, "(b f) d c -> (b d) f c", f=video_length)
if self.pos_encoder is not None:
hidden_states = self.pos_encoder(hidden_states)
hidden_states = rearrange(hidden_states, "(b d) f c -> b (f d) c", f=video_length, d=d)
else:
raise NotImplementedError
encoder_hidden_states = encoder_hidden_states
if self.group_norm is not None:
hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = self.to_q(hidden_states)
dim = query.shape[-1]
query = self.reshape_heads_to_batch_dim(query)
if self.added_kv_proj_dim is not None:
raise NotImplementedError
encoder_hidden_states = encoder_hidden_states if encoder_hidden_states is not None else hidden_states
key = self.to_k(encoder_hidden_states)
value = self.to_v(encoder_hidden_states)
key = self.reshape_heads_to_batch_dim(key)
value = self.reshape_heads_to_batch_dim(value)
if attention_mask is not None:
if attention_mask.shape[-1] != query.shape[1]:
target_length = query.shape[1]
attention_mask = F.pad(attention_mask, (0, target_length), value=0.0)
attention_mask = attention_mask.repeat_interleave(self.heads, dim=0)
bs = 512
new_hidden_states = []
for i in range(0, query.shape[0], bs):
# attention, what we cannot get enough of
if self._use_memory_efficient_attention_xformers:
hidden_states = self._memory_efficient_attention_xformers(query[i : i + bs], key[i : i + bs], value[i : i + bs], attention_mask[i : i + bs] if attention_mask is not None else attention_mask)
# Some versions of xformers return output in fp32, cast it back to the dtype of the input
hidden_states = hidden_states.to(query.dtype)
else:
if self._slice_size is None or query[i : i + bs].shape[0] // self._slice_size == 1:
hidden_states = self._attention(query[i : i + bs], key[i : i + bs], value[i : i + bs], attention_mask[i : i + bs] if attention_mask is not None else attention_mask)
else:
hidden_states = self._sliced_attention(query[i : i + bs], key[i : i + bs], value[i : i + bs], sequence_length, dim, attention_mask[i : i + bs] if attention_mask is not None else attention_mask)
new_hidden_states.append(hidden_states)
for i in range(0, hidden_states.shape[0], bs):
__hidden_states = super().forward(
hidden_states[i : i + bs],
encoder_hidden_states=encoder_hidden_states[i : i + bs],
attention_mask=attention_mask
)
new_hidden_states.append(__hidden_states)
hidden_states = torch.cat(new_hidden_states, dim = 0)
# linear proj
hidden_states = self.to_out[0](hidden_states)
# dropout
hidden_states = self.to_out[1](hidden_states)
if self.attention_mode == "Temporal":
hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d)
if self.grid:
hidden_states = rearrange(hidden_states, "(b f n m) (h w) c -> (b f) h n w m c", f=video_length, n=self.block_size, m=self.block_size, h=height // self.block_size, w=weight // self.block_size)
hidden_states = rearrange(hidden_states, "b h n w m c -> b (h n) (w m) c")
hidden_states = rearrange(hidden_states, "b h w c -> b (h w) c")
elif self.attention_mode == "Global":
hidden_states = rearrange(hidden_states, "b (f d) c -> (b f) d c", f=video_length, d=d)
return hidden_states
+150
View File
@@ -0,0 +1,150 @@
from typing import Any, Dict, Optional, Tuple
import torch
import torch.nn.functional as F
from diffusers.models.embeddings import (CombinedTimestepLabelEmbeddings,
TimestepEmbedding, Timesteps)
from torch import nn
def zero_module(module):
# Zero out the parameters of a module and return it.
for p in module.parameters():
p.detach().zero_()
return module
class FP32LayerNorm(nn.LayerNorm):
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
origin_dtype = inputs.dtype
if hasattr(self, 'weight') and self.weight is not None:
return F.layer_norm(
inputs.float(), self.normalized_shape, self.weight.float(), self.bias.float(), self.eps
).to(origin_dtype)
else:
return F.layer_norm(
inputs.float(), self.normalized_shape, None, None, self.eps
).to(origin_dtype)
class PixArtAlphaCombinedTimestepSizeEmbeddings(nn.Module):
"""
For PixArt-Alpha.
Reference:
https://github.com/PixArt-alpha/PixArt-alpha/blob/0f55e922376d8b797edd44d25d0e7464b260dcab/diffusion/model/nets/PixArtMS.py#L164C9-L168C29
"""
def __init__(self, embedding_dim, size_emb_dim, use_additional_conditions: bool = False):
super().__init__()
self.outdim = size_emb_dim
self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)
self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)
self.use_additional_conditions = use_additional_conditions
if use_additional_conditions:
self.additional_condition_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)
self.resolution_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=size_emb_dim)
self.aspect_ratio_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=size_emb_dim)
self.resolution_embedder.linear_2 = zero_module(self.resolution_embedder.linear_2)
self.aspect_ratio_embedder.linear_2 = zero_module(self.aspect_ratio_embedder.linear_2)
def forward(self, timestep, resolution, aspect_ratio, batch_size, hidden_dtype):
timesteps_proj = self.time_proj(timestep)
timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D)
if self.use_additional_conditions:
resolution_emb = self.additional_condition_proj(resolution.flatten()).to(hidden_dtype)
resolution_emb = self.resolution_embedder(resolution_emb).reshape(batch_size, -1)
aspect_ratio_emb = self.additional_condition_proj(aspect_ratio.flatten()).to(hidden_dtype)
aspect_ratio_emb = self.aspect_ratio_embedder(aspect_ratio_emb).reshape(batch_size, -1)
conditioning = timesteps_emb + torch.cat([resolution_emb, aspect_ratio_emb], dim=1)
else:
conditioning = timesteps_emb
return conditioning
class AdaLayerNormSingle(nn.Module):
r"""
Norm layer adaptive layer norm single (adaLN-single).
As proposed in PixArt-Alpha (see: https://arxiv.org/abs/2310.00426; Section 2.3).
Parameters:
embedding_dim (`int`): The size of each embedding vector.
use_additional_conditions (`bool`): To use additional conditions for normalization or not.
"""
def __init__(self, embedding_dim: int, use_additional_conditions: bool = False):
super().__init__()
self.emb = PixArtAlphaCombinedTimestepSizeEmbeddings(
embedding_dim, size_emb_dim=embedding_dim // 3, use_additional_conditions=use_additional_conditions
)
self.silu = nn.SiLU()
self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=True)
def forward(
self,
timestep: torch.Tensor,
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
batch_size: Optional[int] = None,
hidden_dtype: Optional[torch.dtype] = None,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
# No modulation happening here.
embedded_timestep = self.emb(timestep, **added_cond_kwargs, batch_size=batch_size, hidden_dtype=hidden_dtype)
return self.linear(self.silu(embedded_timestep)), embedded_timestep
class AdaLayerNormShift(nn.Module):
r"""
Norm layer modified to incorporate timestep embeddings.
Parameters:
embedding_dim (`int`): The size of each embedding vector.
num_embeddings (`int`): The size of the embeddings dictionary.
"""
def __init__(self, embedding_dim: int, elementwise_affine=True, eps=1e-6):
super().__init__()
self.silu = nn.SiLU()
self.linear = nn.Linear(embedding_dim, embedding_dim)
self.norm = FP32LayerNorm(embedding_dim, elementwise_affine=elementwise_affine, eps=eps)
def forward(self, x: torch.Tensor, emb: torch.Tensor) -> torch.Tensor:
shift = self.linear(self.silu(emb.to(torch.float32)).to(emb.dtype))
x = self.norm(x) + shift.unsqueeze(dim=1)
return x
class EasyAnimateLayerNormZero(nn.Module):
# Modified from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
# Add fp32 layer norm
def __init__(
self,
conditioning_dim: int,
embedding_dim: int,
elementwise_affine: bool = True,
eps: float = 1e-5,
bias: bool = True,
norm_type: str = "fp32_layer_norm",
) -> None:
super().__init__()
self.silu = nn.SiLU()
self.linear = nn.Linear(conditioning_dim, 6 * embedding_dim, bias=bias)
if norm_type == "layer_norm":
self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=elementwise_affine, eps=eps)
elif norm_type == "fp32_layer_norm":
self.norm = FP32LayerNorm(embedding_dim, elementwise_affine=elementwise_affine, eps=eps)
else:
raise ValueError(
f"Unsupported `norm_type` ({norm_type}) provided. Supported ones are: 'layer_norm', 'fp32_layer_norm'."
)
def forward(
self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
shift, scale, gate, enc_shift, enc_scale, enc_gate = self.linear(self.silu(temb)).chunk(6, dim=1)
hidden_states = self.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :]
encoder_hidden_states = self.norm(encoder_hidden_states) * (1 + enc_scale)[:, None, :] + enc_shift[:, None, :]
return hidden_states, encoder_hidden_states, gate[:, None, :], enc_gate[:, None, :]
+148 -2
View File
@@ -1,3 +1,4 @@
import math
from typing import Optional
import numpy as np
@@ -127,6 +128,86 @@ class UnPatch1D(nn.Module):
return outputs
class Upsampler(nn.Module):
def __init__(
self,
spatial_upsample_factor: int = 1,
temporal_upsample_factor: int = 1,
):
super().__init__()
self.spatial_upsample_factor = spatial_upsample_factor
self.temporal_upsample_factor = temporal_upsample_factor
class TemporalUpsampler3D(Upsampler):
def __init__(self):
super().__init__(
spatial_upsample_factor=1,
temporal_upsample_factor=2,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
if x.shape[2] > 1:
first_frame, x = x[:, :, :1], x[:, :, 1:]
x = F.interpolate(x, scale_factor=(2, 1, 1), mode="trilinear")
x = torch.cat([first_frame, x], dim=2)
return x
class CausalConv3d(nn.Conv3d):
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size=3, # : int | tuple[int, int, int],
stride=1, # : int | tuple[int, int, int] = 1,
padding=1, # : int | tuple[int, int, int], # TODO: change it to 0.
dilation=1, # : int | tuple[int, int, int] = 1,
**kwargs,
):
kernel_size = kernel_size if isinstance(kernel_size, tuple) else (kernel_size,) * 3
assert len(kernel_size) == 3, f"Kernel size must be a 3-tuple, got {kernel_size} instead."
stride = stride if isinstance(stride, tuple) else (stride,) * 3
assert len(stride) == 3, f"Stride must be a 3-tuple, got {stride} instead."
dilation = dilation if isinstance(dilation, tuple) else (dilation,) * 3
assert len(dilation) == 3, f"Dilation must be a 3-tuple, got {dilation} instead."
t_ks, h_ks, w_ks = kernel_size
_, h_stride, w_stride = stride
t_dilation, h_dilation, w_dilation = dilation
t_pad = (t_ks - 1) * t_dilation
# TODO: align with SD
if padding is None:
h_pad = math.ceil(((h_ks - 1) * h_dilation + (1 - h_stride)) / 2)
w_pad = math.ceil(((w_ks - 1) * w_dilation + (1 - w_stride)) / 2)
elif isinstance(padding, int):
h_pad = w_pad = padding
else:
assert NotImplementedError
self.temporal_padding = t_pad
super().__init__(
in_channels=in_channels,
out_channels=out_channels,
kernel_size=kernel_size,
stride=stride,
dilation=dilation,
padding=(0, h_pad, w_pad),
**kwargs,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: (B, C, T, H, W)
x = F.pad(
x,
pad=(0, 0, 0, 0, self.temporal_padding, 0),
mode="replicate", # TODO: check if this is necessary
)
return super().forward(x)
class PatchEmbed3D(nn.Module):
"""3D Image to Patch Embedding"""
@@ -177,7 +258,6 @@ class PatchEmbed3D(nn.Module):
latent = latent.flatten(2).transpose(1, 2) # BCFHW -> BNC
if self.layer_norm:
latent = self.norm(latent)
# Interpolate positional embeddings if needed.
# (For PixArt-Alpha: https://github.com/PixArt-alpha/PixArt-alpha/blob/0f55e922376d8b797edd44d25d0e7464b260dcab/diffusion/model/nets/PixArtMS.py#L162C151-L162C160)
if self.height != height or self.width != width:
@@ -268,4 +348,70 @@ class PatchEmbedF3D(nn.Module):
else:
pos_embed = self.pos_embed
return (latent + pos_embed).to(latent.dtype)
return (latent + pos_embed).to(latent.dtype)
class CasualPatchEmbed3D(nn.Module):
"""3D Image to Patch Embedding"""
def __init__(
self,
height=224,
width=224,
patch_size=16,
time_patch_size=4,
in_channels=3,
embed_dim=768,
layer_norm=False,
flatten=True,
bias=True,
interpolation_scale=1,
):
super().__init__()
num_patches = (height // patch_size) * (width // patch_size)
self.flatten = flatten
self.layer_norm = layer_norm
self.proj = CausalConv3d(
in_channels, embed_dim, kernel_size=(time_patch_size, patch_size, patch_size), stride=(time_patch_size, patch_size, patch_size), bias=bias, padding=None
)
if layer_norm:
self.norm = nn.LayerNorm(embed_dim, elementwise_affine=False, eps=1e-6)
else:
self.norm = None
self.patch_size = patch_size
# See:
# https://github.com/PixArt-alpha/PixArt-alpha/blob/0f55e922376d8b797edd44d25d0e7464b260dcab/diffusion/model/nets/PixArtMS.py#L161
self.height, self.width = height // patch_size, width // patch_size
self.base_size = height // patch_size
self.interpolation_scale = interpolation_scale
pos_embed = get_2d_sincos_pos_embed(
embed_dim, int(num_patches**0.5), base_size=self.base_size, interpolation_scale=self.interpolation_scale
)
self.register_buffer("pos_embed", torch.from_numpy(pos_embed).float().unsqueeze(0), persistent=False)
def forward(self, latent):
height, width = latent.shape[-2] // self.patch_size, latent.shape[-1] // self.patch_size
latent = self.proj(latent)
latent = rearrange(latent, "b c f h w -> (b f) c h w")
if self.flatten:
latent = latent.flatten(2).transpose(1, 2) # BCFHW -> BNC
if self.layer_norm:
latent = self.norm(latent)
# Interpolate positional embeddings if needed.
# (For PixArt-Alpha: https://github.com/PixArt-alpha/PixArt-alpha/blob/0f55e922376d8b797edd44d25d0e7464b260dcab/diffusion/model/nets/PixArtMS.py#L162C151-L162C160)
if self.height != height or self.width != width:
pos_embed = get_2d_sincos_pos_embed(
embed_dim=self.pos_embed.shape[-1],
grid_size=(height, width),
base_size=self.base_size,
interpolation_scale=self.interpolation_scale,
)
pos_embed = torch.from_numpy(pos_embed)
pos_embed = pos_embed.float().unsqueeze(0).to(latent.device)
else:
pos_embed = self.pos_embed
return (latent + pos_embed).to(latent.dtype)
+353
View File
@@ -0,0 +1,353 @@
from typing import Optional
import torch
import torch.nn.functional as F
from diffusers.models.attention import Attention
from diffusers.models.embeddings import apply_rotary_emb
from einops import rearrange, repeat
try:
import xfuser
from xfuser.core.distributed import (
get_sequence_parallel_world_size,
get_sequence_parallel_rank,
get_sp_group,
initialize_model_parallel,
init_distributed_environment
)
from xfuser.core.long_ctx_attention import xFuserLongContextAttention
except Exception as ex:
get_sequence_parallel_world_size = None
get_sequence_parallel_rank = None
xFuserLongContextAttention = None
class HunyuanAttnProcessor2_0:
r"""
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is
used in the HunyuanDiT model. It applies a s normalization layer and rotary embedding on query and key vector.
"""
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
def __call__(
self,
attn: Attention,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
temb: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[torch.Tensor] = None,
) -> torch.Tensor:
residual = hidden_states
if attn.spatial_norm is not None:
hidden_states = attn.spatial_norm(hidden_states, temb)
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
# scaled_dot_product_attention expects attention_mask shape to be
# (batch, heads, source_length, target_length)
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
elif attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)
# Apply RoPE if needed
if image_rotary_emb is not None:
query = apply_rotary_emb(query, image_rotary_emb)
if not attn.is_cross_attention:
key = apply_rotary_emb(key, image_rotary_emb)
# the output of sdp = (batch, num_heads, seq_len, head_dim)
# TODO: add support for attn.scale when we move to Torch 2.1
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
hidden_states = hidden_states.to(query.dtype)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
class LazyKVCompressionProcessor2_0:
r"""
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is
used in the KVCompression model. It applies a s normalization layer and rotary embedding on query and key vector.
"""
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
def __call__(
self,
attn: Attention,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
temb: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[torch.Tensor] = None,
) -> torch.Tensor:
residual = hidden_states
if attn.spatial_norm is not None:
hidden_states = attn.spatial_norm(hidden_states, temb)
input_ndim = hidden_states.ndim
batch_size, channel, num_frames, height, width = hidden_states.shape
hidden_states = rearrange(hidden_states, "b c f h w -> b (f h w) c", f=num_frames, h=height, w=width)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
# scaled_dot_product_attention expects attention_mask shape to be
# (batch, heads, source_length, target_length)
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
elif attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
key = rearrange(key, "b (f h w) c -> (b f) c h w", f=num_frames, h=height, w=width)
key = attn.k_compression(key)
key_shape = key.size()
key = rearrange(key, "(b f) c h w -> b (f h w) c", f=num_frames)
value = rearrange(value, "b (f h w) c -> (b f) c h w", f=num_frames, h=height, w=width)
value = attn.v_compression(value)
value = rearrange(value, "(b f) c h w -> b (f h w) c", f=num_frames)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)
# Apply RoPE if needed
if image_rotary_emb is not None:
compression_image_rotary_emb = (
rearrange(image_rotary_emb[0], "(f h w) c -> f c h w", f=num_frames, h=height, w=width),
rearrange(image_rotary_emb[1], "(f h w) c -> f c h w", f=num_frames, h=height, w=width),
)
compression_image_rotary_emb = (
F.interpolate(compression_image_rotary_emb[0], size=key_shape[-2:], mode='bilinear'),
F.interpolate(compression_image_rotary_emb[1], size=key_shape[-2:], mode='bilinear')
)
compression_image_rotary_emb = (
rearrange(compression_image_rotary_emb[0], "f c h w -> (f h w) c"),
rearrange(compression_image_rotary_emb[1], "f c h w -> (f h w) c"),
)
query = apply_rotary_emb(query, image_rotary_emb)
if not attn.is_cross_attention:
key = apply_rotary_emb(key, compression_image_rotary_emb)
# the output of sdp = (batch, num_heads, seq_len, head_dim)
# TODO: add support for attn.scale when we move to Torch 2.1
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
hidden_states = hidden_states.to(query.dtype)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
class EasyAnimateAttnProcessor2_0:
def __init__(self):
if xFuserLongContextAttention is not None:
try:
get_sequence_parallel_world_size()
self.hybrid_seq_parallel_attn = xFuserLongContextAttention()
except Exception:
self.hybrid_seq_parallel_attn = None
else:
self.hybrid_seq_parallel_attn = None
def __call__(
self,
attn: Attention,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[torch.Tensor] = None,
attn2: Attention = None,
) -> torch.Tensor:
text_seq_length = encoder_hidden_states.size(1)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
if attn2 is None:
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)
if attn2 is not None:
query_txt = attn2.to_q(encoder_hidden_states)
key_txt = attn2.to_k(encoder_hidden_states)
value_txt = attn2.to_v(encoder_hidden_states)
inner_dim = key_txt.shape[-1]
head_dim = inner_dim // attn.heads
query_txt = query_txt.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key_txt = key_txt.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value_txt = value_txt.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
if attn2.norm_q is not None:
query_txt = attn2.norm_q(query_txt)
if attn2.norm_k is not None:
key_txt = attn2.norm_k(key_txt)
query = torch.cat([query_txt, query], dim=2)
key = torch.cat([key_txt, key], dim=2)
value = torch.cat([value_txt, value], dim=2)
# Apply RoPE if needed
if image_rotary_emb is not None:
query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
if not attn.is_cross_attention:
key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
if self.hybrid_seq_parallel_attn is None:
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2)
else:
sp_world_rank = get_sequence_parallel_rank()
sp_world_size = get_sequence_parallel_world_size()
img_q = query[:, :, text_seq_length:].transpose(1,2)
txt_q = query[:, :, :text_seq_length].transpose(1,2)
img_k = key[:, :, text_seq_length:].transpose(1,2)
txt_k = key[:, :, :text_seq_length].transpose(1,2)
img_v = value[:, :, text_seq_length:].transpose(1,2)
txt_v = value[:, :, :text_seq_length].transpose(1,2)
hidden_states = self.hybrid_seq_parallel_attn(None,
img_q, img_k, img_v, dropout_p=0.0, causal=False,
joint_tensor_query=txt_q,
joint_tensor_key=txt_k,
joint_tensor_value=txt_v,
joint_strategy='front',)
hidden_states = hidden_states.reshape(batch_size, -1, attn.heads * head_dim)
if attn2 is None:
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
encoder_hidden_states, hidden_states = hidden_states.split(
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
)
else:
encoder_hidden_states, hidden_states = hidden_states.split(
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
encoder_hidden_states = attn2.to_out[0](encoder_hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
encoder_hidden_states = attn2.to_out[1](encoder_hidden_states)
return hidden_states, encoder_hidden_states
+146
View File
@@ -0,0 +1,146 @@
# Copyright (c) Alibaba Cloud.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
import math
import numpy as np
import torch
from torch import nn
from torch.nn import functional as F
from torch.nn.init import normal_
def get_abs_pos(abs_pos, tgt_size):
# abs_pos: L, C
# tgt_size: M
# return: M, C
src_size = int(math.sqrt(abs_pos.size(0)))
tgt_size = int(math.sqrt(tgt_size))
dtype = abs_pos.dtype
if src_size != tgt_size:
return F.interpolate(
abs_pos.float().reshape(1, src_size, src_size, -1).permute(0, 3, 1, 2),
size=(tgt_size, tgt_size),
mode="bicubic",
align_corners=False,
).permute(0, 2, 3, 1).flatten(0, 2).to(dtype=dtype)
else:
return abs_pos
# https://github.com/facebookresearch/mae/blob/efb2a8062c206524e35e47d04501ed4f544c0ae8/util/pos_embed.py#L20
def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False):
"""
grid_size: int of the grid height and width
return:
pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)
"""
grid_h = np.arange(grid_size, dtype=np.float32)
grid_w = np.arange(grid_size, dtype=np.float32)
grid = np.meshgrid(grid_w, grid_h) # here w goes first
grid = np.stack(grid, axis=0)
grid = grid.reshape([2, 1, grid_size, grid_size])
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
if cls_token:
pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0)
return pos_embed
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
assert embed_dim % 2 == 0
# use half of dimensions to encode grid_h
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
return emb
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
"""
embed_dim: output dimension for each position
pos: a list of positions to be encoded: size (M,)
out: (M, D)
"""
assert embed_dim % 2 == 0
omega = np.arange(embed_dim // 2, dtype=np.float32)
omega /= embed_dim / 2.
omega = 1. / 10000**omega # (D/2,)
pos = pos.reshape(-1) # (M,)
out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
emb_sin = np.sin(out) # (M, D/2)
emb_cos = np.cos(out) # (M, D/2)
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
return emb
class Resampler(nn.Module):
"""
A 2D perceiver-resampler network with one cross attention layers by
(grid_size**2) learnable queries and 2d sincos pos_emb
Outputs:
A tensor with the shape of (grid_size**2, embed_dim)
"""
def __init__(
self,
grid_size,
embed_dim,
num_heads,
kv_dim=None,
norm_layer=nn.LayerNorm
):
super().__init__()
self.num_queries = grid_size ** 2
self.embed_dim = embed_dim
self.num_heads = num_heads
self.pos_embed = nn.Parameter(
torch.from_numpy(get_2d_sincos_pos_embed(embed_dim, grid_size)).float()
).requires_grad_(False)
self.query = nn.Parameter(torch.zeros(self.num_queries, embed_dim))
normal_(self.query, std=.02)
if kv_dim is not None and kv_dim != embed_dim:
self.kv_proj = nn.Linear(kv_dim, embed_dim, bias=False)
else:
self.kv_proj = nn.Identity()
self.attn = nn.MultiheadAttention(embed_dim, num_heads)
self.ln_q = norm_layer(embed_dim)
self.ln_kv = norm_layer(embed_dim)
self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
normal_(m.weight, std=.02)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
def forward(self, x, key_padding_mask=None):
pos_embed = get_abs_pos(self.pos_embed, x.size(1))
x = self.kv_proj(x)
x = self.ln_kv(x).permute(1, 0, 2)
N = x.shape[1]
q = self.ln_q(self.query)
out = self.attn(
self._repeat(q, N) + self.pos_embed.unsqueeze(1),
x + pos_embed.unsqueeze(1),
x,
key_padding_mask=key_padding_mask)[0]
return out.permute(1, 0, 2)
def _repeat(self, query, N: int):
return query.unsqueeze(1).repeat(1, N, 1)
+11
View File
@@ -107,11 +107,14 @@ class Transformer2DModel(ModelMixin, ConfigMixin):
norm_eps: float = 1e-5,
attention_type: str = "default",
caption_channels: int = None,
# block type
basic_block_type: str = "basic",
):
super().__init__()
self.use_linear_projection = use_linear_projection
self.num_attention_heads = num_attention_heads
self.attention_head_dim = attention_head_dim
self.basic_block_type = basic_block_type
inner_dim = num_attention_heads * attention_head_dim
conv_cls = nn.Conv2d if USE_PEFT_BACKEND else LoRACompatibleConv
@@ -375,6 +378,9 @@ class Transformer2DModel(ModelMixin, ConfigMixin):
for block in self.transformer_blocks:
if self.training and self.gradient_checkpointing:
args = {
"basic": [],
}[self.basic_block_type]
hidden_states = torch.utils.checkpoint.checkpoint(
block,
hidden_states,
@@ -384,9 +390,13 @@ class Transformer2DModel(ModelMixin, ConfigMixin):
timestep,
cross_attention_kwargs,
class_labels,
*args,
use_reentrant=False,
)
else:
kwargs = {
"basic": {},
}[self.basic_block_type]
hidden_states = block(
hidden_states,
attention_mask=attention_mask,
@@ -395,6 +405,7 @@ class Transformer2DModel(ModelMixin, ConfigMixin):
timestep=timestep,
cross_attention_kwargs=cross_attention_kwargs,
class_labels=class_labels,
**kwargs
)
# 3. Output
File diff suppressed because it is too large Load Diff
+43 -40
View File
@@ -12,6 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
import copy
import html
import inspect
import re
@@ -153,7 +154,7 @@ class EasyAnimatePipeline(DiffusionPipeline):
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
# Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/utils.py
def mask_text_embeddings(self, emb, mask):
if emb.shape[0] == 1:
@@ -520,7 +521,9 @@ class EasyAnimatePipeline(DiffusionPipeline):
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_latents
def prepare_latents(self, batch_size, num_channels_latents, video_length, height, width, dtype, device, generator, latents=None):
if self.vae.quant_conv.weight.ndim==5:
shape = (batch_size, num_channels_latents, int(video_length // 21 * 6), height // self.vae_scale_factor, width // self.vae_scale_factor)
mini_batch_encoder = self.vae.mini_batch_encoder
mini_batch_decoder = self.vae.mini_batch_decoder
shape = (batch_size, num_channels_latents, int(video_length // mini_batch_encoder * mini_batch_decoder) if video_length != 1 else 1, height // self.vae_scale_factor, width // self.vae_scale_factor)
else:
shape = (batch_size, num_channels_latents, video_length, height // self.vae_scale_factor, width // self.vae_scale_factor)
@@ -538,46 +541,34 @@ class EasyAnimatePipeline(DiffusionPipeline):
# scale the initial noise by the standard deviation required by the scheduler
latents = latents * self.scheduler.init_noise_sigma
return latents
def smooth_output(self, video, mini_batch_encoder, mini_batch_decoder):
if video.size()[2] <= mini_batch_encoder:
return video
prefix_index_before = mini_batch_encoder // 2
prefix_index_after = mini_batch_encoder - prefix_index_before
pixel_values = video[:, :, prefix_index_before:-prefix_index_after]
# Encode middle videos
latents = self.vae.encode(pixel_values)[0]
latents = latents.mode()
# Decode middle videos
middle_video = self.vae.decode(latents)[0]
video[:, :, prefix_index_before:-prefix_index_after] = (video[:, :, prefix_index_before:-prefix_index_after] + middle_video) / 2
return video
def decode_latents(self, latents):
video_length = latents.shape[2]
latents = 1 / 0.18215 * latents
latents = 1 / self.vae.config.scaling_factor * latents
if self.vae.quant_conv.weight.ndim==5:
video = []
mini_batch = 6
latents = latents.float()
self.vae.post_quant_conv = self.vae.post_quant_conv.float()
self.vae.decoder = self.vae.decoder.float()
for i in range(0, latents.shape[2], mini_batch):
with torch.no_grad():
start_index = i
end_index = i + mini_batch
latents_bs = self.vae.decode(latents[:, :, start_index:end_index, :, :])[0]
video.append(latents_bs)
# for i in range(0, latents.shape[2], mini_batch // 4):
# if i == latents.shape[2] - 1:
# break
# with torch.no_grad():
# if i == 0:
# start_index = i
# end_index = i + mini_batch // 4 + 1
# latents_bs = self.vae.decode(latents[:, :, start_index:end_index, :, :])[0]
# else:
# start_index = i + 1
# end_index = i + mini_batch // 4 + 1
# weight_dtype = self.vae.encoder.encoder[0].conv.weight.dtype
# previous_video_encoded = self.vae.encode(video[-1][:, :, -1:].to(weight_dtype))[0]
# previous_video_encoded = previous_video_encoded.sample()
# latents_bs = self.vae.decode(
# torch.cat([previous_video_encoded, latents[:, :, start_index:end_index, :, :], ], 2)
# )[0][:, :, 1:, :, :]
video = torch.cat(video, 2)
mini_batch_encoder = self.vae.mini_batch_encoder
mini_batch_decoder = self.vae.mini_batch_decoder
video = self.vae.decode(latents)[0]
video = video.clamp(-1, 1)
video = self.smooth_output(video, mini_batch_encoder, mini_batch_decoder).cpu().clamp(-1, 1)
else:
latents = rearrange(latents, "b c f h w -> (b f) c h w")
# video = self.vae.decode(latents).sample
video = []
for frame_idx in tqdm(range(latents.shape[0])):
video.append(self.vae.decode(latents[frame_idx:frame_idx+1]).sample)
@@ -614,6 +605,7 @@ class EasyAnimatePipeline(DiffusionPipeline):
callback_steps: int = 1,
clean_caption: bool = True,
max_sequence_length: int = 120,
comfyui_progressbar: bool = False,
**kwargs,
) -> Union[EasyAnimatePipelineOutput, Tuple]:
"""
@@ -698,6 +690,12 @@ class EasyAnimatePipeline(DiffusionPipeline):
batch_size = prompt_embeds.shape[0]
device = self._execution_device
if self.text_encoder is not None:
dtype = self.text_encoder.dtype
elif self.text_encoder_2 is not None:
dtype = self.text_encoder_2.dtype
else:
dtype = self.transformer.dtype
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
@@ -738,7 +736,7 @@ class EasyAnimatePipeline(DiffusionPipeline):
video_length,
height,
width,
prompt_embeds.dtype,
dtype,
device,
generator,
latents,
@@ -752,8 +750,8 @@ class EasyAnimatePipeline(DiffusionPipeline):
if self.transformer.config.sample_size == 128:
resolution = torch.tensor([height, width]).repeat(batch_size * num_images_per_prompt, 1)
aspect_ratio = torch.tensor([float(height / width)]).repeat(batch_size * num_images_per_prompt, 1)
resolution = resolution.to(dtype=prompt_embeds.dtype, device=device)
aspect_ratio = aspect_ratio.to(dtype=prompt_embeds.dtype, device=device)
resolution = resolution.to(dtype=dtype, device=device)
aspect_ratio = aspect_ratio.to(dtype=dtype, device=device)
if do_classifier_free_guidance:
resolution = torch.cat([resolution, resolution], dim=0)
@@ -763,7 +761,9 @@ class EasyAnimatePipeline(DiffusionPipeline):
# 7. Denoising loop
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
if comfyui_progressbar:
from comfy.utils import ProgressBar
pbar = ProgressBar(num_inference_steps)
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
@@ -815,6 +815,9 @@ class EasyAnimatePipeline(DiffusionPipeline):
step_idx = i // getattr(self.scheduler, "order", 1)
callback(step_idx, t, latents)
if comfyui_progressbar:
pbar.update(1)
# Post-processing
video = self.decode_latents(latents)
@@ -12,6 +12,8 @@
# See the License for the specific language governing permissions and
# limitations under the License.
import copy
import gc
import html
import inspect
import re
@@ -21,6 +23,7 @@ from typing import Callable, List, Optional, Tuple, Union
import numpy as np
import torch
import torch.nn.functional as F
from diffusers import DiffusionPipeline, ImagePipelineOutput
from diffusers.image_processor import VaeImageProcessor
from diffusers.models import AutoencoderKL
@@ -30,8 +33,10 @@ from diffusers.utils import (BACKENDS_MAPPING, BaseOutput, deprecate,
replace_example_docstring)
from diffusers.utils.torch_utils import randn_tensor
from einops import rearrange
from PIL import Image
from tqdm import tqdm
from transformers import T5EncoderModel, T5Tokenizer
from transformers import (CLIPImageProcessor, CLIPVisionModelWithProjection,
T5EncoderModel, T5Tokenizer)
from ..models.transformer3d import Transformer3DModel
@@ -108,11 +113,15 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
vae: AutoencoderKL,
transformer: Transformer3DModel,
scheduler: DPMSolverMultistepScheduler,
clip_image_processor:CLIPImageProcessor = None,
clip_image_encoder:CLIPVisionModelWithProjection = None,
):
super().__init__()
self.register_modules(
tokenizer=tokenizer, text_encoder=text_encoder, vae=vae, transformer=transformer, scheduler=scheduler
tokenizer=tokenizer, text_encoder=text_encoder, vae=vae, transformer=transformer,
scheduler=scheduler,
clip_image_processor=clip_image_processor, clip_image_encoder=clip_image_encoder,
)
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
@@ -140,8 +149,11 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
device: Optional[torch.device] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
prompt_attention_mask: Optional[torch.FloatTensor] = None,
negative_prompt_attention_mask: Optional[torch.FloatTensor] = None,
clean_caption: bool = False,
mask_feature: bool = True,
max_sequence_length: int = 120,
**kwargs,
):
r"""
Encodes the prompt into text encoder hidden states.
@@ -165,11 +177,15 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
negative_prompt_embeds (`torch.FloatTensor`, *optional*):
Pre-generated negative text embeddings. For PixArt-Alpha, it's should be the embeddings of the ""
string.
clean_caption (bool, defaults to `False`):
clean_caption (`bool`, defaults to `False`):
If `True`, the function will preprocess and clean the provided caption before encoding.
mask_feature: (bool, defaults to `True`):
If `True`, the function will mask the text embeddings.
max_sequence_length (`int`, defaults to 120): Maximum sequence length to use for the prompt.
"""
if "mask_feature" in kwargs:
deprecation_message = "The use of `mask_feature` is deprecated. It is no longer used in any computation and that doesn't affect the end results. It will be removed in a future version."
deprecate("mask_feature", "1.0.0", deprecation_message, standard_warn=False)
if device is None:
device = self._execution_device
@@ -181,7 +197,7 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
batch_size = prompt_embeds.shape[0]
# See Section 3.1. of the paper.
max_length = 120
max_length = max_sequence_length
if prompt_embeds is None:
prompt = self._text_preprocessing(prompt, clean_caption=clean_caption)
@@ -205,13 +221,11 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
f" {max_length} tokens: {removed_text}"
)
attention_mask = text_inputs.attention_mask.to(device)
prompt_embeds_attention_mask = attention_mask
prompt_attention_mask = text_inputs.attention_mask
prompt_attention_mask = prompt_attention_mask.to(device)
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=attention_mask)
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask)
prompt_embeds = prompt_embeds[0]
else:
prompt_embeds_attention_mask = torch.ones_like(prompt_embeds)
if self.text_encoder is not None:
dtype = self.text_encoder.dtype
@@ -226,8 +240,8 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
# duplicate text embeddings and attention mask for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1)
prompt_embeds_attention_mask = prompt_embeds_attention_mask.view(bs_embed, -1)
prompt_embeds_attention_mask = prompt_embeds_attention_mask.repeat(num_images_per_prompt, 1)
prompt_attention_mask = prompt_attention_mask.view(bs_embed, -1)
prompt_attention_mask = prompt_attention_mask.repeat(num_images_per_prompt, 1)
# get unconditional embeddings for classifier free guidance
if do_classifier_free_guidance and negative_prompt_embeds is None:
@@ -243,11 +257,11 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
add_special_tokens=True,
return_tensors="pt",
)
attention_mask = uncond_input.attention_mask.to(device)
negative_prompt_attention_mask = uncond_input.attention_mask
negative_prompt_attention_mask = negative_prompt_attention_mask.to(device)
negative_prompt_embeds = self.text_encoder(
uncond_input.input_ids.to(device),
attention_mask=attention_mask,
uncond_input.input_ids.to(device), attention_mask=negative_prompt_attention_mask
)
negative_prompt_embeds = negative_prompt_embeds[0]
@@ -260,23 +274,13 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
negative_prompt_embeds = negative_prompt_embeds.repeat(1, num_images_per_prompt, 1)
negative_prompt_embeds = negative_prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1)
# For classifier free guidance, we need to do two forward passes.
# Here we concatenate the unconditional and text embeddings into a single batch
# to avoid doing two forward passes
negative_prompt_attention_mask = negative_prompt_attention_mask.view(bs_embed, -1)
negative_prompt_attention_mask = negative_prompt_attention_mask.repeat(num_images_per_prompt, 1)
else:
negative_prompt_embeds = None
negative_prompt_attention_mask = None
# Perform additional masking.
if mask_feature:
prompt_embeds = prompt_embeds.unsqueeze(1)
masked_prompt_embeds, keep_indices = self.mask_text_embeddings(prompt_embeds, prompt_embeds_attention_mask)
masked_prompt_embeds = masked_prompt_embeds.squeeze(1)
masked_negative_prompt_embeds = (
negative_prompt_embeds[:, :keep_indices, :] if negative_prompt_embeds is not None else None
)
return masked_prompt_embeds, masked_negative_prompt_embeds
return prompt_embeds, negative_prompt_embeds
return prompt_embeds, prompt_attention_mask, negative_prompt_embeds, negative_prompt_attention_mask
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_extra_step_kwargs
def prepare_extra_step_kwargs(self, generator, eta):
@@ -489,6 +493,60 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
return caption.strip()
def prepare_mask_latents(
self, mask, masked_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance
):
# resize the mask to latents shape as we concatenate the mask to the latents
# we do that before converting to dtype to avoid breaking in case we're using cpu_offload
# and half precision
video_length = mask.shape[2]
mask = mask.to(device=device, dtype=self.vae.dtype)
if self.vae.quant_conv.weight.ndim==5:
bs = 1
new_mask = []
for i in range(0, mask.shape[0], bs):
mask_bs = mask[i : i + bs]
mask_bs = self.vae.encode(mask_bs)[0]
mask_bs = mask_bs.sample()
new_mask.append(mask_bs)
mask = torch.cat(new_mask, dim = 0)
mask = mask * self.vae.config.scaling_factor
else:
if mask.shape[1] == 4:
mask = mask
else:
video_length = mask.shape[2]
mask = rearrange(mask, "b c f h w -> (b f) c h w")
mask = self._encode_vae_image(mask, generator=generator)
mask = rearrange(mask, "(b f) c h w -> b c f h w", f=video_length)
masked_image = masked_image.to(device=device, dtype=self.vae.dtype)
if self.vae.quant_conv.weight.ndim==5:
bs = 1
new_mask_pixel_values = []
for i in range(0, masked_image.shape[0], bs):
mask_pixel_values_bs = masked_image[i : i + bs]
mask_pixel_values_bs = self.vae.encode(mask_pixel_values_bs)[0]
mask_pixel_values_bs = mask_pixel_values_bs.sample()
new_mask_pixel_values.append(mask_pixel_values_bs)
masked_image_latents = torch.cat(new_mask_pixel_values, dim = 0)
masked_image_latents = masked_image_latents * self.vae.config.scaling_factor
else:
if masked_image.shape[1] == 4:
masked_image_latents = masked_image
else:
video_length = mask.shape[2]
masked_image = rearrange(masked_image, "b c f h w -> (b f) c h w")
masked_image_latents = self._encode_vae_image(masked_image, generator=generator)
masked_image_latents = rearrange(masked_image_latents, "(b f) c h w -> b c f h w", f=video_length)
# aligning device to prevent device errors when concating it with the latent model input
masked_image_latents = masked_image_latents.to(device=device, dtype=dtype)
return mask, masked_image_latents
def prepare_latents(
self,
batch_size,
@@ -506,39 +564,54 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
return_noise=False,
return_video_latents=False,
):
shape = (batch_size, num_channels_latents, video_length, height // self.vae_scale_factor, width // self.vae_scale_factor)
if self.vae.quant_conv.weight.ndim==5:
mini_batch_encoder = self.vae.mini_batch_encoder
mini_batch_decoder = self.vae.mini_batch_decoder
shape = (batch_size, num_channels_latents, int(video_length // mini_batch_encoder * mini_batch_decoder) if video_length != 1 else 1, height // self.vae_scale_factor, width // self.vae_scale_factor)
else:
shape = (batch_size, num_channels_latents, video_length, height // self.vae_scale_factor, width // self.vae_scale_factor)
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
if return_video_latents or (latents is None and not is_strength_max):
video = video.to(device=device, dtype=dtype)
if video.shape[1] == 4:
video_latents = video
if return_video_latents or (latents is None and not is_strength_max):
video = video.to(device=device, dtype=self.vae.dtype)
if self.vae.quant_conv.weight.ndim==5:
bs = 1
mini_batch_encoder = self.vae.mini_batch_encoder
new_video = []
for i in range(0, video.shape[0], bs):
video_bs = video[i : i + bs]
video_bs = self.vae.encode(video_bs)[0]
video_bs = video_bs.sample()
new_video.append(video_bs)
video = torch.cat(new_video, dim = 0)
video = video * self.vae.config.scaling_factor
else:
video_length = video.shape[2]
video = rearrange(video, "b c f h w -> (b f) c h w")
video_latents = self._encode_vae_image(image=video, generator=generator)
video_latents = rearrange(video_latents, "(b f) c h w -> b c f h w", f=video_length)
video_latents = video_latents.repeat(batch_size // video_latents.shape[0], 1, 1, 1, 1)
if video.shape[1] == 4:
video = video
else:
video_length = video.shape[2]
video = rearrange(video, "b c f h w -> (b f) c h w")
video = self._encode_vae_image(video, generator=generator)
video = rearrange(video, "(b f) c h w -> b c f h w", f=video_length)
video_latents = video.repeat(batch_size // video.shape[0], 1, 1, 1, 1)
if latents is None:
rand_device = "cpu" if device.type == "mps" else device
noise = torch.randn(shape, generator=generator, device=rand_device, dtype=dtype).to(device)
noise = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
# if strength is 1. then initialise the latents to noise, else initial to image + noise
latents = noise if is_strength_max else self.scheduler.add_noise(video_latents, noise, timestep)
# if pure noise then scale the initial latents by the Scheduler's init sigma
latents = latents * self.scheduler.init_noise_sigma if is_strength_max else latents
else:
noise = latents.to(device)
if latents.shape != shape:
raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {shape}")
latents = latents.to(device)
latents = noise * self.scheduler.init_noise_sigma
# scale the initial noise by the standard deviation required by the scheduler
latents = latents * self.scheduler.init_noise_sigma
outputs = (latents,)
if return_noise:
@@ -549,16 +622,38 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
return outputs
def smooth_output(self, video, mini_batch_encoder, mini_batch_decoder):
if video.size()[2] <= mini_batch_encoder:
return video
prefix_index_before = mini_batch_encoder // 2
prefix_index_after = mini_batch_encoder - prefix_index_before
pixel_values = video[:, :, prefix_index_before:-prefix_index_after]
# Encode middle videos
latents = self.vae.encode(pixel_values)[0]
latents = latents.sample()
# Decode middle videos
middle_video = self.vae.decode(latents)[0]
video[:, :, prefix_index_before:-prefix_index_after] = (video[:, :, prefix_index_before:-prefix_index_after] + middle_video) / 2
return video
def decode_latents(self, latents):
video_length = latents.shape[2]
latents = 1 / 0.18215 * latents
latents = rearrange(latents, "b c f h w -> (b f) c h w")
# video = self.vae.decode(latents).sample
video = []
for frame_idx in tqdm(range(latents.shape[0])):
video.append(self.vae.decode(latents[frame_idx:frame_idx+1]).sample)
video = torch.cat(video)
video = rearrange(video, "(b f) c h w -> b c f h w", f=video_length)
latents = 1 / self.vae.config.scaling_factor * latents
if self.vae.quant_conv.weight.ndim==5:
mini_batch_encoder = self.vae.mini_batch_encoder
mini_batch_decoder = self.vae.mini_batch_decoder
video = self.vae.decode(latents)[0]
video = video.clamp(-1, 1)
video = self.smooth_output(video, mini_batch_encoder, mini_batch_decoder).cpu().clamp(-1, 1)
else:
latents = rearrange(latents, "b c f h w -> (b f) c h w")
video = []
for frame_idx in tqdm(range(latents.shape[0])):
video.append(self.vae.decode(latents[frame_idx:frame_idx+1]).sample)
video = torch.cat(video)
video = rearrange(video, "(b f) c h w -> b c f h w", f=video_length)
video = (video / 2 + 0.5).clamp(0, 1)
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16
video = video.cpu().float().numpy()
@@ -578,39 +673,16 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
return image_latents
def prepare_mask_latents(
self, mask, masked_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance
):
# resize the mask to latents shape as we concatenate the mask to the latents
# we do that before converting to dtype to avoid breaking in case we're using cpu_offload
# and half precision
video_length = mask.shape[2]
mask = rearrange(mask, "b c f h w -> (b f) c h w")
mask = torch.nn.functional.interpolate(
mask, size=(height // self.vae_scale_factor, width // self.vae_scale_factor)
)
mask = mask.to(device=device, dtype=dtype)
mask = rearrange(mask, "(b f) c h w -> b c f h w", f=video_length)
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.StableDiffusionImg2ImgPipeline.get_timesteps
def get_timesteps(self, num_inference_steps, strength, device):
# get the original timestep using init_timestep
init_timestep = min(int(num_inference_steps * strength), num_inference_steps)
masked_image = masked_image.to(device=device, dtype=dtype)
t_start = max(num_inference_steps - init_timestep, 0)
timesteps = self.scheduler.timesteps[t_start * self.scheduler.order :]
if masked_image.shape[1] == 4:
masked_image_latents = masked_image
else:
video_length = mask.shape[2]
masked_image = rearrange(masked_image, "b c f h w -> (b f) c h w")
masked_image_latents = self._encode_vae_image(masked_image, generator=generator)
masked_image_latents = rearrange(masked_image_latents, "(b f) c h w -> b c f h w", f=video_length)
return timesteps, num_inference_steps - t_start
mask = torch.cat([mask] * 2) if do_classifier_free_guidance else mask
masked_image_latents = (
torch.cat([masked_image_latents] * 2) if do_classifier_free_guidance else masked_image_latents
)
# aligning device to prevent device errors when concating it with the latent model input
masked_image_latents = masked_image_latents.to(device=device, dtype=dtype)
return mask, masked_image_latents
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(
@@ -632,13 +704,20 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
prompt_attention_mask: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: Optional[str] = "tensor",
negative_prompt_attention_mask: Optional[torch.FloatTensor] = None,
output_type: Optional[str] = "latent",
return_dict: bool = True,
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
callback_steps: int = 1,
clean_caption: bool = True,
mask_feature: bool = True,
max_sequence_length: int = 120,
clip_image: Image = None,
clip_apply_ratio: float = 0.50,
comfyui_progressbar: bool = False,
**kwargs,
) -> Union[EasyAnimatePipelineOutput, Tuple]:
"""
Function invoked when calling the pipeline for generation.
@@ -712,6 +791,8 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
# 1. Check inputs. Raise error if not correct
height = height or self.transformer.config.sample_size * self.vae_scale_factor
width = width or self.transformer.config.sample_size * self.vae_scale_factor
height = int(height // 16 * 16)
width = int(width // 16 * 16)
# 2. Default height and width to transformer
if prompt is not None and isinstance(prompt, str):
@@ -722,6 +803,12 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
batch_size = prompt_embeds.shape[0]
device = self._execution_device
if self.text_encoder is not None:
dtype = self.text_encoder.dtype
elif self.text_encoder_2 is not None:
dtype = self.text_encoder_2.dtype
else:
dtype = self.transformer.dtype
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
@@ -729,7 +816,12 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
do_classifier_free_guidance = guidance_scale > 1.0
# 3. Encode input prompt
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
(
prompt_embeds,
prompt_attention_mask,
negative_prompt_embeds,
negative_prompt_attention_mask,
) = self.encode_prompt(
prompt,
do_classifier_free_guidance,
negative_prompt=negative_prompt,
@@ -737,29 +829,37 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
device=device,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
prompt_attention_mask=prompt_attention_mask,
negative_prompt_attention_mask=negative_prompt_attention_mask,
clean_caption=clean_caption,
mask_feature=mask_feature,
max_sequence_length=max_sequence_length,
)
if do_classifier_free_guidance:
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask], dim=0)
# 4. Prepare timesteps
# 4. set timesteps
self.scheduler.set_timesteps(num_inference_steps, device=device)
timesteps = self.scheduler.timesteps
timesteps, num_inference_steps = self.get_timesteps(
num_inference_steps=num_inference_steps, strength=strength, device=device
)
# at which timestep to set the initial noise (n.b. 50% if strength is 0.5)
latent_timestep = timesteps[:1].repeat(batch_size)
latent_timestep = timesteps[:1].repeat(batch_size * num_images_per_prompt)
# create a boolean to check if the strength is set to 1. if so then initialise the latents with pure noise
is_strength_max = strength == 1.0
video_length = video.shape[2]
init_video = self.image_processor.preprocess(rearrange(video, "b c f h w -> (b f) c h w"), height=height, width=width)
init_video = init_video.to(dtype=torch.float32)
init_video = rearrange(init_video, "(b f) c h w -> b c f h w", f=video_length)
if video is not None:
video_length = video.shape[2]
init_video = self.image_processor.preprocess(rearrange(video, "b c f h w -> (b f) c h w"), height=height, width=width)
init_video = init_video.to(dtype=torch.float32)
init_video = rearrange(init_video, "(b f) c h w -> b c f h w", f=video_length)
else:
init_video = None
# Prepare latent variables
num_channels_latents = self.vae.config.latent_channels
num_channels_transformer = self.transformer.config.in_channels
return_image_latents = num_channels_transformer == 4
return_image_latents = True # num_channels_transformer == 4
# 5. Prepare latents.
latents_outputs = self.prepare_latents(
@@ -768,7 +868,7 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
height,
width,
video_length,
prompt_embeds.dtype,
dtype,
device,
generator,
latents,
@@ -782,34 +882,92 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
latents, noise, image_latents = latents_outputs
else:
latents, noise = latents_outputs
latents_dtype = latents.dtype
latents_dtype = dtype
# Prepare mask latent variables
video_length = video.shape[2]
mask_condition = self.mask_processor.preprocess(rearrange(mask_video, "b c f h w -> (b f) c h w"), height=height, width=width)
mask_condition = mask_condition.to(dtype=torch.float32)
mask_condition = rearrange(mask_condition, "(b f) c h w -> b c f h w", f=video_length)
if mask_video is not None:
# Prepare mask latent variables
video_length = video.shape[2]
mask_condition = self.mask_processor.preprocess(rearrange(mask_video, "b c f h w -> (b f) c h w"), height=height, width=width)
mask_condition = mask_condition.to(dtype=torch.float32)
mask_condition = rearrange(mask_condition, "(b f) c h w -> b c f h w", f=video_length)
if masked_video_latents is None:
masked_video = init_video * (mask_condition < 0.5)
if num_channels_transformer == 12:
mask_condition_tile = torch.tile(mask_condition, [1, 3, 1, 1, 1])
if masked_video_latents is None:
masked_video = init_video * (mask_condition_tile < 0.5) + torch.ones_like(init_video) * (mask_condition_tile > 0.5) * -1
else:
masked_video = masked_video_latents
mask_latents, masked_video_latents = self.prepare_mask_latents(
mask_condition_tile,
masked_video,
batch_size,
height,
width,
dtype,
device,
generator,
do_classifier_free_guidance,
)
mask = torch.tile(mask_condition, [1, num_channels_transformer // 3, 1, 1, 1])
mask = F.interpolate(mask, size=latents.size()[-3:], mode='trilinear', align_corners=True).to(device, dtype)
mask_input = torch.cat([mask_latents] * 2) if do_classifier_free_guidance else mask_latents
masked_video_latents_input = (
torch.cat([masked_video_latents] * 2) if do_classifier_free_guidance else masked_video_latents
)
inpaint_latents = torch.cat([mask_input, masked_video_latents_input], dim=1).to(dtype)
else:
mask = torch.tile(mask_condition, [1, num_channels_transformer, 1, 1, 1])
mask = F.interpolate(mask, size=latents.size()[-3:], mode='trilinear', align_corners=True).to(device, dtype)
inpaint_latents = None
else:
masked_video = masked_video_latents
if num_channels_transformer == 12:
mask = torch.zeros_like(latents).to(device, dtype)
masked_video_latents = torch.zeros_like(latents).to(device, dtype)
mask_input = torch.cat([mask] * 2) if do_classifier_free_guidance else mask
masked_video_latents_input = (
torch.cat([masked_video_latents] * 2) if do_classifier_free_guidance else masked_video_latents
)
inpaint_latents = torch.cat([mask_input, masked_video_latents_input], dim=1).to(dtype)
else:
mask = torch.zeros_like(init_video[:, :1])
mask = torch.tile(mask, [1, num_channels_transformer, 1, 1, 1])
mask = F.interpolate(mask, size=latents.size()[-3:], mode='trilinear', align_corners=True).to(device, dtype)
inpaint_latents = None
if clip_image is not None:
inputs = self.clip_image_processor(images=clip_image, return_tensors="pt")
inputs["pixel_values"] = inputs["pixel_values"].to(device, dtype=dtype)
clip_encoder_hidden_states = self.clip_image_encoder(**inputs).image_embeds
clip_encoder_hidden_states_neg = torch.zeros([batch_size, 768]).to(device, dtype=dtype)
clip_attention_mask = torch.ones([batch_size, 8]).to(device, dtype=dtype)
clip_attention_mask_neg = torch.zeros([batch_size, 8]).to(device, dtype=dtype)
clip_encoder_hidden_states_input = torch.cat([clip_encoder_hidden_states_neg, clip_encoder_hidden_states]) if do_classifier_free_guidance else clip_encoder_hidden_states
clip_attention_mask_input = torch.cat([clip_attention_mask_neg, clip_attention_mask]) if do_classifier_free_guidance else clip_attention_mask
elif clip_image is None and num_channels_transformer == 12:
clip_encoder_hidden_states = torch.zeros([batch_size, 768]).to(device, dtype=dtype)
clip_attention_mask = torch.zeros([batch_size, 8])
clip_attention_mask = clip_attention_mask.to(device, dtype=dtype)
clip_encoder_hidden_states_input = torch.cat([clip_encoder_hidden_states] * 2) if do_classifier_free_guidance else clip_encoder_hidden_states
clip_attention_mask_input = torch.cat([clip_attention_mask] * 2) if do_classifier_free_guidance else clip_attention_mask
else:
clip_encoder_hidden_states_input = None
clip_attention_mask_input = None
mask, masked_video_latents = self.prepare_mask_latents(
mask_condition,
masked_video,
batch_size,
height,
width,
prompt_embeds.dtype,
device,
generator,
do_classifier_free_guidance,
)
# Check that sizes of mask, masked image and latents match
if num_channels_transformer == 9:
if num_channels_transformer == 12:
# default case for runwayml/stable-diffusion-inpainting
num_channels_mask = mask.shape[1]
num_channels_mask = mask_latents.shape[1]
num_channels_masked_image = masked_video_latents.shape[1]
if num_channels_latents + num_channels_mask + num_channels_masked_image != self.transformer.config.in_channels:
raise ValueError(
@@ -819,12 +977,12 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
f" = {num_channels_latents+num_channels_masked_image+num_channels_mask}. Please verify the config of"
" `pipeline.transformer` or your `mask_image` or `image` input."
)
elif num_channels_transformer == 4:
elif num_channels_transformer != 4:
raise ValueError(
f"The transformer {self.transformer.__class__} should have 9 input channels, not {self.transformer.config.in_channels}."
)
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
# 9. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
# 6.1 Prepare micro-conditions.
@@ -832,20 +990,32 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
if self.transformer.config.sample_size == 128:
resolution = torch.tensor([height, width]).repeat(batch_size * num_images_per_prompt, 1)
aspect_ratio = torch.tensor([float(height / width)]).repeat(batch_size * num_images_per_prompt, 1)
resolution = resolution.to(dtype=prompt_embeds.dtype, device=device)
aspect_ratio = aspect_ratio.to(dtype=prompt_embeds.dtype, device=device)
resolution = resolution.to(dtype=dtype, device=device)
aspect_ratio = aspect_ratio.to(dtype=dtype, device=device)
if do_classifier_free_guidance:
resolution = torch.cat([resolution, resolution], dim=0)
aspect_ratio = torch.cat([aspect_ratio, aspect_ratio], dim=0)
added_cond_kwargs = {"resolution": resolution, "aspect_ratio": aspect_ratio}
# 7. Denoising loop
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
# 10. Denoising loop
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
self._num_timesteps = len(timesteps)
if comfyui_progressbar:
from comfy.utils import ProgressBar
pbar = ProgressBar(num_inference_steps)
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
if num_channels_transformer == 9:
latent_model_input = torch.cat([latent_model_input, mask, masked_video_latents], dim=1)
if i < len(timesteps) * (1 - clip_apply_ratio) and clip_encoder_hidden_states_input is not None:
clip_encoder_hidden_states_actual_input = torch.zeros_like(clip_encoder_hidden_states_input)
clip_attention_mask_actual_input = torch.zeros_like(clip_attention_mask_input)
else:
clip_encoder_hidden_states_actual_input = clip_encoder_hidden_states_input
clip_attention_mask_actual_input = clip_attention_mask_input
current_timestep = t
if not torch.is_tensor(current_timestep):
@@ -866,8 +1036,12 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
noise_pred = self.transformer(
latent_model_input,
encoder_hidden_states=prompt_embeds,
encoder_attention_mask=prompt_attention_mask,
timestep=current_timestep,
added_cond_kwargs=added_cond_kwargs,
inpaint_latents=inpaint_latents,
clip_encoder_hidden_states=clip_encoder_hidden_states_actual_input,
clip_attention_mask=clip_attention_mask_actual_input,
return_dict=False,
)[0]
@@ -882,6 +1056,17 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
# compute previous image: x_t -> x_t-1
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
if num_channels_transformer == 4:
init_latents_proper = image_latents
init_mask = mask
if i < len(timesteps) - 1:
noise_timestep = timesteps[i + 1]
init_latents_proper = self.scheduler.add_noise(
init_latents_proper, noise, torch.tensor([noise_timestep])
)
latents = (1 - init_mask) * init_latents_proper + init_mask * latents
# call the callback, if provided
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
@@ -889,11 +1074,14 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
step_idx = i // getattr(self.scheduler, "order", 1)
callback(step_idx, t, latents)
if comfyui_progressbar:
pbar.update(1)
# Post-processing
video = self.decode_latents(latents)
# Convert to tensor
if output_type == "tensor":
if output_type == "latent":
video = torch.from_numpy(video)
if not return_dict:
@@ -0,0 +1,918 @@
# Copyright 2024 EasyAnimate Authors and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import inspect
from typing import Callable, Dict, List, Optional, Tuple, Union
import numpy as np
import torch
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.image_processor import VaeImageProcessor
from diffusers.models.embeddings import (get_2d_rotary_pos_embed,
get_3d_rotary_pos_embed)
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput
from diffusers.pipelines.stable_diffusion.safety_checker import \
StableDiffusionSafetyChecker
from diffusers.schedulers import DDIMScheduler
from diffusers.utils import (is_torch_xla_available, logging,
replace_example_docstring)
from diffusers.utils.torch_utils import randn_tensor
from einops import rearrange
from tqdm import tqdm
from transformers import (BertModel, BertTokenizer, CLIPImageProcessor,
T5Tokenizer, T5EncoderModel)
from .pipeline_easyanimate import EasyAnimatePipelineOutput
from ..models import AutoencoderKLMagvit, EasyAnimateTransformer3DModel
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
XLA_AVAILABLE = True
else:
XLA_AVAILABLE = False
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
EXAMPLE_DOC_STRING = """
Examples:
```py
>>> pass
```
"""
def get_resize_crop_region_for_grid(src, tgt_width, tgt_height):
tw = tgt_width
th = tgt_height
h, w = src
r = h / w
if r > (th / tw):
resize_height = th
resize_width = int(round(th / h * w))
else:
resize_width = tw
resize_height = int(round(tw / w * h))
crop_top = int(round((th - resize_height) / 2.0))
crop_left = int(round((tw - resize_width) / 2.0))
return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width)
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.rescale_noise_cfg
def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0):
"""
Rescale `noise_cfg` according to `guidance_rescale`. Based on findings of [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). See Section 3.4
"""
std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True)
std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)
# rescale the results from guidance (fixes overexposure)
noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
# mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images
noise_cfg = guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg
return noise_cfg
class EasyAnimatePipeline_Multi_Text_Encoder(DiffusionPipeline):
r"""
Pipeline for text-to-video generation using EasyAnimate.
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods the
library implements for all the pipelines (such as downloading or saving, running on a particular device, etc.)
EasyAnimate uses two text encoders: [mT5](https://huggingface.co/google/mt5-base) and [bilingual CLIP](fine-tuned by
HunyuanDiT team)
Args:
vae ([`AutoencoderKLMagvit`]):
Variational Auto-Encoder (VAE) Model to encode and decode video to and from latent representations.
text_encoder (Optional[`~transformers.BertModel`, `~transformers.CLIPTextModel`]):
Frozen text-encoder ([clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14)).
EasyAnimate uses a fine-tuned [bilingual CLIP].
tokenizer (Optional[`~transformers.BertTokenizer`, `~transformers.CLIPTokenizer`]):
A `BertTokenizer` or `CLIPTokenizer` to tokenize text.
transformer ([`EasyAnimateTransformer3DModel`]):
The EasyAnimate model designed by Tencent Hunyuan.
text_encoder_2 (`T5EncoderModel`):
The mT5 embedder.
tokenizer_2 (`T5Tokenizer`):
The tokenizer for the mT5 embedder.
scheduler ([`DDIMScheduler`]):
A scheduler to be used in combination with EasyAnimate to denoise the encoded image latents.
"""
model_cpu_offload_seq = "text_encoder->text_encoder_2->transformer->vae"
_optional_components = [
"safety_checker",
"feature_extractor",
"text_encoder_2",
"tokenizer_2",
"text_encoder",
"tokenizer",
]
_exclude_from_cpu_offload = ["safety_checker"]
_callback_tensor_inputs = [
"latents",
"prompt_embeds",
"negative_prompt_embeds",
"prompt_embeds_2",
"negative_prompt_embeds_2",
]
def __init__(
self,
vae: AutoencoderKLMagvit,
text_encoder: BertModel,
tokenizer: BertTokenizer,
text_encoder_2: T5EncoderModel,
tokenizer_2: T5Tokenizer,
transformer: EasyAnimateTransformer3DModel,
scheduler: DDIMScheduler,
safety_checker: StableDiffusionSafetyChecker,
feature_extractor: CLIPImageProcessor,
requires_safety_checker: bool = True,
):
super().__init__()
self.register_modules(
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
tokenizer_2=tokenizer_2,
transformer=transformer,
scheduler=scheduler,
safety_checker=safety_checker,
feature_extractor=feature_extractor,
text_encoder_2=text_encoder_2,
)
if safety_checker is None and requires_safety_checker:
logger.warning(
f"You have disabled the safety checker for {self.__class__} by passing `safety_checker=None`. Ensure"
" that you abide to the conditions of the Stable Diffusion license and do not expose unfiltered"
" results in services or applications open to the public. Both the diffusers team and Hugging Face"
" strongly recommend to keep the safety filter enabled in all public facing circumstances, disabling"
" it only for use-cases that involve analyzing network behavior or auditing its results. For more"
" information, please have a look at https://github.com/huggingface/diffusers/pull/254 ."
)
if safety_checker is not None and feature_extractor is None:
raise ValueError(
"Make sure to define a feature extractor when loading {self.__class__} if you want to use the safety"
" checker. If you do not want to use the safety checker, you can pass `'safety_checker=None'` instead."
)
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
self.register_to_config(requires_safety_checker=requires_safety_checker)
def enable_sequential_cpu_offload(self, *args, **kwargs):
super().enable_sequential_cpu_offload(*args, **kwargs)
if hasattr(self.transformer, "clip_projection") and self.transformer.clip_projection is not None:
import accelerate
accelerate.hooks.remove_hook_from_module(self.transformer.clip_projection, recurse=True)
self.transformer.clip_projection = self.transformer.clip_projection.to("cuda")
def encode_prompt(
self,
prompt: str,
device: torch.device,
dtype: torch.dtype,
num_images_per_prompt: int = 1,
do_classifier_free_guidance: bool = True,
negative_prompt: Optional[str] = None,
prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
prompt_attention_mask: Optional[torch.Tensor] = None,
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
max_sequence_length: Optional[int] = None,
text_encoder_index: int = 0,
actual_max_sequence_length: int = 256
):
r"""
Encodes the prompt into text encoder hidden states.
Args:
prompt (`str` or `List[str]`, *optional*):
prompt to be encoded
device: (`torch.device`):
torch device
dtype (`torch.dtype`):
torch dtype
num_images_per_prompt (`int`):
number of images that should be generated per prompt
do_classifier_free_guidance (`bool`):
whether to use classifier free guidance or not
negative_prompt (`str` or `List[str]`, *optional*):
The prompt or prompts not to guide the image generation. If not defined, one has to pass
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
less than `1`).
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
provided, text embeddings will be generated from `prompt` input argument.
negative_prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
argument.
prompt_attention_mask (`torch.Tensor`, *optional*):
Attention mask for the prompt. Required when `prompt_embeds` is passed directly.
negative_prompt_attention_mask (`torch.Tensor`, *optional*):
Attention mask for the negative prompt. Required when `negative_prompt_embeds` is passed directly.
max_sequence_length (`int`, *optional*): maximum sequence length to use for the prompt.
text_encoder_index (`int`, *optional*):
Index of the text encoder to use. `0` for clip and `1` for T5.
"""
tokenizers = [self.tokenizer, self.tokenizer_2]
text_encoders = [self.text_encoder, self.text_encoder_2]
tokenizer = tokenizers[text_encoder_index]
text_encoder = text_encoders[text_encoder_index]
if max_sequence_length is None:
if text_encoder_index == 0:
max_length = min(self.tokenizer.model_max_length, actual_max_sequence_length)
if text_encoder_index == 1:
max_length = min(self.tokenizer_2.model_max_length, actual_max_sequence_length)
else:
max_length = max_sequence_length
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
if prompt_embeds is None:
text_inputs = tokenizer(
prompt,
padding="max_length",
max_length=max_length,
truncation=True,
return_attention_mask=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
if text_input_ids.shape[-1] > actual_max_sequence_length:
reprompt = tokenizer.batch_decode(text_input_ids[:, :actual_max_sequence_length], skip_special_tokens=True)
text_inputs = tokenizer(
reprompt,
padding="max_length",
max_length=max_length,
truncation=True,
return_attention_mask=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
untruncated_ids = tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(
text_input_ids, untruncated_ids
):
_actual_max_sequence_length = min(tokenizer.model_max_length, actual_max_sequence_length)
removed_text = tokenizer.batch_decode(untruncated_ids[:, _actual_max_sequence_length - 1 : -1])
logger.warning(
"The following part of your input was truncated because CLIP can only handle sequences up to"
f" {_actual_max_sequence_length} tokens: {removed_text}"
)
prompt_attention_mask = text_inputs.attention_mask.to(device)
if self.transformer.config.enable_text_attention_mask:
prompt_embeds = text_encoder(
text_input_ids.to(device),
attention_mask=prompt_attention_mask,
)
else:
prompt_embeds = text_encoder(
text_input_ids.to(device)
)
prompt_embeds = prompt_embeds[0]
prompt_attention_mask = prompt_attention_mask.repeat(num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
bs_embed, seq_len, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1)
# get unconditional embeddings for classifier free guidance
if do_classifier_free_guidance and negative_prompt_embeds is None:
uncond_tokens: List[str]
if negative_prompt is None:
uncond_tokens = [""] * batch_size
elif prompt is not None and type(prompt) is not type(negative_prompt):
raise TypeError(
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
f" {type(prompt)}."
)
elif isinstance(negative_prompt, str):
uncond_tokens = [negative_prompt]
elif batch_size != len(negative_prompt):
raise ValueError(
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
" the batch size of `prompt`."
)
else:
uncond_tokens = negative_prompt
max_length = prompt_embeds.shape[1]
uncond_input = tokenizer(
uncond_tokens,
padding="max_length",
max_length=max_length,
truncation=True,
return_tensors="pt",
)
uncond_input_ids = uncond_input.input_ids
if uncond_input_ids.shape[-1] > actual_max_sequence_length:
reuncond_tokens = tokenizer.batch_decode(uncond_input_ids[:, :actual_max_sequence_length], skip_special_tokens=True)
uncond_input = tokenizer(
reuncond_tokens,
padding="max_length",
max_length=max_length,
truncation=True,
return_attention_mask=True,
return_tensors="pt",
)
uncond_input_ids = uncond_input.input_ids
negative_prompt_attention_mask = uncond_input.attention_mask.to(device)
if self.transformer.config.enable_text_attention_mask:
negative_prompt_embeds = text_encoder(
uncond_input.input_ids.to(device),
attention_mask=negative_prompt_attention_mask,
)
else:
negative_prompt_embeds = text_encoder(
uncond_input.input_ids.to(device)
)
negative_prompt_embeds = negative_prompt_embeds[0]
negative_prompt_attention_mask = negative_prompt_attention_mask.repeat(num_images_per_prompt, 1)
if do_classifier_free_guidance:
# duplicate unconditional embeddings for each generation per prompt, using mps friendly method
seq_len = negative_prompt_embeds.shape[1]
negative_prompt_embeds = negative_prompt_embeds.to(dtype=dtype, device=device)
negative_prompt_embeds = negative_prompt_embeds.repeat(1, num_images_per_prompt, 1)
negative_prompt_embeds = negative_prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1)
return prompt_embeds, negative_prompt_embeds, prompt_attention_mask, negative_prompt_attention_mask
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.run_safety_checker
def run_safety_checker(self, image, device, dtype):
if self.safety_checker is None:
has_nsfw_concept = None
else:
if torch.is_tensor(image):
feature_extractor_input = self.image_processor.postprocess(image, output_type="pil")
else:
feature_extractor_input = self.image_processor.numpy_to_pil(image)
safety_checker_input = self.feature_extractor(feature_extractor_input, return_tensors="pt").to(device)
image, has_nsfw_concept = self.safety_checker(
images=image, clip_input=safety_checker_input.pixel_values.to(dtype)
)
return image, has_nsfw_concept
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_extra_step_kwargs
def prepare_extra_step_kwargs(self, generator, eta):
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
# eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.
# eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502
# and should be between [0, 1]
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
extra_step_kwargs = {}
if accepts_eta:
extra_step_kwargs["eta"] = eta
# check if the scheduler accepts generator
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
if accepts_generator:
extra_step_kwargs["generator"] = generator
return extra_step_kwargs
def check_inputs(
self,
prompt,
height,
width,
negative_prompt=None,
prompt_embeds=None,
negative_prompt_embeds=None,
prompt_attention_mask=None,
negative_prompt_attention_mask=None,
prompt_embeds_2=None,
negative_prompt_embeds_2=None,
prompt_attention_mask_2=None,
negative_prompt_attention_mask_2=None,
callback_on_step_end_tensor_inputs=None,
):
if height % 8 != 0 or width % 8 != 0:
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
if prompt is not None and prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
" only forward one of the two."
)
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is None and prompt_embeds_2 is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds_2`. Cannot leave both `prompt` and `prompt_embeds_2` undefined."
)
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
if prompt_embeds is not None and prompt_attention_mask is None:
raise ValueError("Must provide `prompt_attention_mask` when specifying `prompt_embeds`.")
if prompt_embeds_2 is not None and prompt_attention_mask_2 is None:
raise ValueError("Must provide `prompt_attention_mask_2` when specifying `prompt_embeds_2`.")
if negative_prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
)
if negative_prompt_embeds is not None and negative_prompt_attention_mask is None:
raise ValueError("Must provide `negative_prompt_attention_mask` when specifying `negative_prompt_embeds`.")
if negative_prompt_embeds_2 is not None and negative_prompt_attention_mask_2 is None:
raise ValueError(
"Must provide `negative_prompt_attention_mask_2` when specifying `negative_prompt_embeds_2`."
)
if prompt_embeds is not None and negative_prompt_embeds is not None:
if prompt_embeds.shape != negative_prompt_embeds.shape:
raise ValueError(
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
f" {negative_prompt_embeds.shape}."
)
if prompt_embeds_2 is not None and negative_prompt_embeds_2 is not None:
if prompt_embeds_2.shape != negative_prompt_embeds_2.shape:
raise ValueError(
"`prompt_embeds_2` and `negative_prompt_embeds_2` must have the same shape when passed directly, but"
f" got: `prompt_embeds_2` {prompt_embeds_2.shape} != `negative_prompt_embeds_2`"
f" {negative_prompt_embeds_2.shape}."
)
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_latents
def prepare_latents(self, batch_size, num_channels_latents, video_length, height, width, dtype, device, generator, latents=None):
if self.vae.quant_conv is None or self.vae.quant_conv.weight.ndim==5:
if self.vae.cache_mag_vae:
mini_batch_encoder = self.vae.mini_batch_encoder
mini_batch_decoder = self.vae.mini_batch_decoder
shape = (batch_size, num_channels_latents, int((video_length - 1) // mini_batch_encoder * mini_batch_decoder + 1) if video_length != 1 else 1, height // self.vae_scale_factor, width // self.vae_scale_factor)
else:
mini_batch_encoder = self.vae.mini_batch_encoder
mini_batch_decoder = self.vae.mini_batch_decoder
shape = (batch_size, num_channels_latents, int(video_length // mini_batch_encoder * mini_batch_decoder) if video_length != 1 else 1, height // self.vae_scale_factor, width // self.vae_scale_factor)
else:
shape = (batch_size, num_channels_latents, video_length, height // self.vae_scale_factor, width // self.vae_scale_factor)
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
if latents is None:
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
else:
latents = latents.to(device)
# scale the initial noise by the standard deviation required by the scheduler
latents = latents * self.scheduler.init_noise_sigma
return latents
def smooth_output(self, video, mini_batch_encoder, mini_batch_decoder):
if video.size()[2] <= mini_batch_encoder:
return video
prefix_index_before = mini_batch_encoder // 2
prefix_index_after = mini_batch_encoder - prefix_index_before
pixel_values = video[:, :, prefix_index_before:-prefix_index_after]
# Encode middle videos
latents = self.vae.encode(pixel_values)[0]
latents = latents.mode()
# Decode middle videos
middle_video = self.vae.decode(latents)[0]
video[:, :, prefix_index_before:-prefix_index_after] = (video[:, :, prefix_index_before:-prefix_index_after] + middle_video) / 2
return video
def decode_latents(self, latents):
video_length = latents.shape[2]
latents = 1 / self.vae.config.scaling_factor * latents
if self.vae.quant_conv is None or self.vae.quant_conv.weight.ndim==5:
mini_batch_encoder = self.vae.mini_batch_encoder
mini_batch_decoder = self.vae.mini_batch_decoder
video = self.vae.decode(latents)[0]
video = video.clamp(-1, 1)
if not self.vae.cache_compression_vae and not self.vae.cache_mag_vae:
video = self.smooth_output(video, mini_batch_encoder, mini_batch_decoder).cpu().clamp(-1, 1)
else:
latents = rearrange(latents, "b c f h w -> (b f) c h w")
video = []
for frame_idx in tqdm(range(latents.shape[0])):
video.append(self.vae.decode(latents[frame_idx:frame_idx+1]).sample)
video = torch.cat(video)
video = rearrange(video, "(b f) c h w -> b c f h w", f=video_length)
video = (video / 2 + 0.5).clamp(0, 1)
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16
video = video.cpu().float().numpy()
return video
@property
def guidance_scale(self):
return self._guidance_scale
@property
def guidance_rescale(self):
return self._guidance_rescale
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
# corresponds to doing no classifier free guidance.
@property
def do_classifier_free_guidance(self):
return self._guidance_scale > 1
@property
def num_timesteps(self):
return self._num_timesteps
@property
def interrupt(self):
return self._interrupt
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(
self,
prompt: Union[str, List[str]] = None,
video_length: Optional[int] = None,
height: Optional[int] = None,
width: Optional[int] = None,
num_inference_steps: Optional[int] = 50,
guidance_scale: Optional[float] = 5.0,
negative_prompt: Optional[Union[str, List[str]]] = None,
num_images_per_prompt: Optional[int] = 1,
eta: Optional[float] = 0.0,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
prompt_embeds: Optional[torch.Tensor] = None,
prompt_embeds_2: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds_2: Optional[torch.Tensor] = None,
prompt_attention_mask: Optional[torch.Tensor] = None,
prompt_attention_mask_2: Optional[torch.Tensor] = None,
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
negative_prompt_attention_mask_2: Optional[torch.Tensor] = None,
output_type: Optional[str] = "latent",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
guidance_rescale: float = 0.0,
original_size: Optional[Tuple[int, int]] = (1024, 1024),
target_size: Optional[Tuple[int, int]] = None,
crops_coords_top_left: Tuple[int, int] = (0, 0),
comfyui_progressbar: bool = False,
):
r"""
Generates images or video using the EasyAnimate pipeline based on the provided prompts.
Examples:
prompt (`str` or `List[str]`, *optional*):
Text prompts to guide the image or video generation. If not provided, use `prompt_embeds` instead.
video_length (`int`, *optional*):
Length of the generated video (in frames).
height (`int`, *optional*):
Height of the generated image in pixels.
width (`int`, *optional*):
Width of the generated image in pixels.
num_inference_steps (`int`, *optional*, defaults to 50):
Number of denoising steps during generation. More steps generally yield higher quality images but slow down inference.
guidance_scale (`float`, *optional*, defaults to 5.0):
Encourages the model to align outputs with prompts. A higher value may decrease image quality.
negative_prompt (`str` or `List[str]`, *optional*):
Prompts indicating what to exclude in generation. If not specified, use `negative_prompt_embeds`.
num_images_per_prompt (`int`, *optional*, defaults to 1):
Number of images to generate for each prompt.
eta (`float`, *optional*, defaults to 0.0):
Applies to DDIM scheduling. Controlled by the eta parameter from the related literature.
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
A generator to ensure reproducibility in image generation.
latents (`torch.Tensor`, *optional*):
Predefined latent tensors to condition generation.
prompt_embeds (`torch.Tensor`, *optional*):
Text embeddings for the prompts. Overrides prompt string inputs for more flexibility.
prompt_embeds_2 (`torch.Tensor`, *optional*):
Secondary text embeddings to supplement or replace the initial prompt embeddings.
negative_prompt_embeds (`torch.Tensor`, *optional*):
Embeddings for negative prompts. Overrides string inputs if defined.
negative_prompt_embeds_2 (`torch.Tensor`, *optional*):
Secondary embeddings for negative prompts, similar to `negative_prompt_embeds`.
prompt_attention_mask (`torch.Tensor`, *optional*):
Attention mask for the primary prompt embeddings.
prompt_attention_mask_2 (`torch.Tensor`, *optional*):
Attention mask for the secondary prompt embeddings.
negative_prompt_attention_mask (`torch.Tensor`, *optional*):
Attention mask for negative prompt embeddings.
negative_prompt_attention_mask_2 (`torch.Tensor`, *optional*):
Attention mask for secondary negative prompt embeddings.
output_type (`str`, *optional*, defaults to "latent"):
Format of the generated output, either as a PIL image or as a NumPy array.
return_dict (`bool`, *optional*, defaults to `True`):
If `True`, returns a structured output. Otherwise returns a simple tuple.
callback_on_step_end (`Callable`, *optional*):
Functions called at the end of each denoising step.
callback_on_step_end_tensor_inputs (`List[str]`, *optional*):
Tensor names to be included in callback function calls.
guidance_rescale (`float`, *optional*, defaults to 0.0):
Adjusts noise levels based on guidance scale.
original_size (`Tuple[int, int]`, *optional*, defaults to `(1024, 1024)`):
Original dimensions of the output.
target_size (`Tuple[int, int]`, *optional*):
Desired output dimensions for calculations.
crops_coords_top_left (`Tuple[int, int]`, *optional*, defaults to `(0, 0)`):
Coordinates for cropping.
Returns:
[`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] or `tuple`:
If `return_dict` is `True`, [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] is returned,
otherwise a `tuple` is returned where the first element is a list with the generated images and the
second element is a list of `bool`s indicating whether the corresponding generated image contains
"not-safe-for-work" (nsfw) content.
"""
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
# 0. default height and width
height = int((height // 16) * 16)
width = int((width // 16) * 16)
# 1. Check inputs. Raise error if not correct
self.check_inputs(
prompt,
height,
width,
negative_prompt,
prompt_embeds,
negative_prompt_embeds,
prompt_attention_mask,
negative_prompt_attention_mask,
prompt_embeds_2,
negative_prompt_embeds_2,
prompt_attention_mask_2,
negative_prompt_attention_mask_2,
callback_on_step_end_tensor_inputs,
)
self._guidance_scale = guidance_scale
self._guidance_rescale = guidance_rescale
self._interrupt = False
# 2. Define call parameters
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
device = self._execution_device
if self.text_encoder is not None:
dtype = self.text_encoder.dtype
elif self.text_encoder_2 is not None:
dtype = self.text_encoder_2.dtype
else:
dtype = self.transformer.dtype
# 3. Encode input prompt
(
prompt_embeds,
negative_prompt_embeds,
prompt_attention_mask,
negative_prompt_attention_mask,
) = self.encode_prompt(
prompt=prompt,
device=device,
dtype=dtype,
num_images_per_prompt=num_images_per_prompt,
do_classifier_free_guidance=self.do_classifier_free_guidance,
negative_prompt=negative_prompt,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
prompt_attention_mask=prompt_attention_mask,
negative_prompt_attention_mask=negative_prompt_attention_mask,
text_encoder_index=0,
)
(
prompt_embeds_2,
negative_prompt_embeds_2,
prompt_attention_mask_2,
negative_prompt_attention_mask_2,
) = self.encode_prompt(
prompt=prompt,
device=device,
dtype=dtype,
num_images_per_prompt=num_images_per_prompt,
do_classifier_free_guidance=self.do_classifier_free_guidance,
negative_prompt=negative_prompt,
prompt_embeds=prompt_embeds_2,
negative_prompt_embeds=negative_prompt_embeds_2,
prompt_attention_mask=prompt_attention_mask_2,
negative_prompt_attention_mask=negative_prompt_attention_mask_2,
text_encoder_index=1,
)
# 4. Prepare timesteps
self.scheduler.set_timesteps(num_inference_steps, device=device)
timesteps = self.scheduler.timesteps
if comfyui_progressbar:
from comfy.utils import ProgressBar
pbar = ProgressBar(num_inference_steps + 1)
# 5. Prepare latent variables
num_channels_latents = self.transformer.config.in_channels
latents = self.prepare_latents(
batch_size * num_images_per_prompt,
num_channels_latents,
video_length,
height,
width,
dtype,
device,
generator,
latents,
)
if comfyui_progressbar:
pbar.update(1)
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
# 7 create image_rotary_emb, style embedding & time ids
grid_height = height // 8 // self.transformer.config.patch_size
grid_width = width // 8 // self.transformer.config.patch_size
if self.transformer.config.get("time_position_encoding_type", "2d_rope") == "3d_rope":
base_size_width = 720 // 8 // self.transformer.config.patch_size
base_size_height = 480 // 8 // self.transformer.config.patch_size
grid_crops_coords = get_resize_crop_region_for_grid(
(grid_height, grid_width), base_size_width, base_size_height
)
image_rotary_emb = get_3d_rotary_pos_embed(
self.transformer.config.attention_head_dim, grid_crops_coords, grid_size=(grid_height, grid_width),
temporal_size=latents.size(2), use_real=True,
)
else:
base_size = 512 // 8 // self.transformer.config.patch_size
grid_crops_coords = get_resize_crop_region_for_grid(
(grid_height, grid_width), base_size, base_size
)
image_rotary_emb = get_2d_rotary_pos_embed(
self.transformer.config.attention_head_dim, grid_crops_coords, (grid_height, grid_width)
)
# Get other hunyuan params
style = torch.tensor([0], device=device)
target_size = target_size or (height, width)
add_time_ids = list(original_size + target_size + crops_coords_top_left)
add_time_ids = torch.tensor([add_time_ids], dtype=dtype)
if self.do_classifier_free_guidance:
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds])
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask])
prompt_embeds_2 = torch.cat([negative_prompt_embeds_2, prompt_embeds_2])
prompt_attention_mask_2 = torch.cat([negative_prompt_attention_mask_2, prompt_attention_mask_2])
add_time_ids = torch.cat([add_time_ids] * 2, dim=0)
style = torch.cat([style] * 2, dim=0)
# To latents.device
prompt_embeds = prompt_embeds.to(device=device)
prompt_attention_mask = prompt_attention_mask.to(device=device)
prompt_embeds_2 = prompt_embeds_2.to(device=device)
prompt_attention_mask_2 = prompt_attention_mask_2.to(device=device)
add_time_ids = add_time_ids.to(dtype=dtype, device=device).repeat(
batch_size * num_images_per_prompt, 1
)
style = style.to(device=device).repeat(batch_size * num_images_per_prompt)
# 8. Denoising loop
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
self._num_timesteps = len(timesteps)
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if self.interrupt:
continue
# expand the latents if we are doing classifier free guidance
latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
# expand scalar t to 1-D tensor to match the 1st dim of latent_model_input
t_expand = torch.tensor([t] * latent_model_input.shape[0], device=device).to(
dtype=latent_model_input.dtype
)
# predict the noise residual
noise_pred = self.transformer(
latent_model_input,
t_expand,
encoder_hidden_states=prompt_embeds,
text_embedding_mask=prompt_attention_mask,
encoder_hidden_states_t5=prompt_embeds_2,
text_embedding_mask_t5=prompt_attention_mask_2,
image_meta_size=add_time_ids,
style=style,
image_rotary_emb=image_rotary_emb,
return_dict=False,
)[0]
if noise_pred.size()[1] != self.vae.config.latent_channels:
noise_pred, _ = noise_pred.chunk(2, dim=1)
# perform guidance
if self.do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
if self.do_classifier_free_guidance and guidance_rescale > 0.0:
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=guidance_rescale)
# compute the previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
if callback_on_step_end is not None:
callback_kwargs = {}
for k in callback_on_step_end_tensor_inputs:
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
prompt_embeds_2 = callback_outputs.pop("prompt_embeds_2", prompt_embeds_2)
negative_prompt_embeds_2 = callback_outputs.pop(
"negative_prompt_embeds_2", negative_prompt_embeds_2
)
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if XLA_AVAILABLE:
xm.mark_step()
if comfyui_progressbar:
pbar.update(1)
# Post-processing
video = self.decode_latents(latents)
# Convert to tensor
if output_type == "latent":
video = torch.from_numpy(video)
# Offload all models
self.maybe_free_model_hooks()
if not return_dict:
return video
return EasyAnimatePipelineOutput(videos=video)
@@ -0,0 +1,989 @@
# Copyright 2024 EasyAnimate Authors and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import inspect
import re
import urllib.parse as ul
from dataclasses import dataclass
from typing import Callable, Dict, List, Optional, Tuple, Union
import numpy as np
import torch
import torch.nn.functional as F
from diffusers import DiffusionPipeline, ImagePipelineOutput
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.image_processor import VaeImageProcessor
from diffusers.models import AutoencoderKL, HunyuanDiT2DModel
from diffusers.models.embeddings import (get_2d_rotary_pos_embed,
get_3d_rotary_pos_embed)
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput
from diffusers.pipelines.stable_diffusion.safety_checker import \
StableDiffusionSafetyChecker
from diffusers.schedulers import DDIMScheduler, DPMSolverMultistepScheduler
from diffusers.utils import (BACKENDS_MAPPING, BaseOutput, deprecate,
is_bs4_available, is_ftfy_available,
is_torch_xla_available, logging,
replace_example_docstring)
from diffusers.utils.torch_utils import randn_tensor
from einops import rearrange
from PIL import Image
from tqdm import tqdm
from transformers import (BertModel, BertTokenizer, CLIPImageProcessor,
CLIPVisionModelWithProjection,
T5EncoderModel, T5Tokenizer)
from ..models import AutoencoderKLMagvit, EasyAnimateTransformer3DModel
from .pipeline_easyanimate import EasyAnimatePipelineOutput
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
XLA_AVAILABLE = True
else:
XLA_AVAILABLE = False
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
EXAMPLE_DOC_STRING = """
Examples:
```py
>>> pass
```
"""
def get_resize_crop_region_for_grid(src, tgt_width, tgt_height):
tw = tgt_width
th = tgt_height
h, w = src
r = h / w
if r > (th / tw):
resize_height = th
resize_width = int(round(th / h * w))
else:
resize_width = tw
resize_height = int(round(tw / w * h))
crop_top = int(round((th - resize_height) / 2.0))
crop_left = int(round((tw - resize_width) / 2.0))
return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width)
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.rescale_noise_cfg
def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0):
"""
Rescale `noise_cfg` according to `guidance_rescale`. Based on findings of [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). See Section 3.4
"""
std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True)
std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)
# rescale the results from guidance (fixes overexposure)
noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
# mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images
noise_cfg = guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg
return noise_cfg
class EasyAnimatePipeline_Multi_Text_Encoder_Control(DiffusionPipeline):
r"""
Pipeline for text-to-video generation using EasyAnimate.
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods the
library implements for all the pipelines (such as downloading or saving, running on a particular device, etc.)
EasyAnimate uses two text encoders: [mT5](https://huggingface.co/google/mt5-base) and [bilingual CLIP](fine-tuned by
HunyuanDiT team)
Args:
vae ([`AutoencoderKLMagvit`]):
Variational Auto-Encoder (VAE) Model to encode and decode video to and from latent representations.
text_encoder (Optional[`~transformers.BertModel`, `~transformers.CLIPTextModel`]):
Frozen text-encoder ([clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14)).
EasyAnimate uses a fine-tuned [bilingual CLIP].
tokenizer (Optional[`~transformers.BertTokenizer`, `~transformers.CLIPTokenizer`]):
A `BertTokenizer` or `CLIPTokenizer` to tokenize text.
transformer ([`EasyAnimateTransformer3DModel`]):
The EasyAnimate model designed by Tencent Hunyuan.
text_encoder_2 (`T5EncoderModel`):
The mT5 embedder.
tokenizer_2 (`T5Tokenizer`):
The tokenizer for the mT5 embedder.
scheduler ([`DDIMScheduler`]):
A scheduler to be used in combination with EasyAnimate to denoise the encoded image latents.
"""
model_cpu_offload_seq = "text_encoder->text_encoder_2->transformer->vae"
_optional_components = [
"safety_checker",
"feature_extractor",
"text_encoder_2",
"tokenizer_2",
"text_encoder",
"tokenizer",
]
_exclude_from_cpu_offload = ["safety_checker"]
_callback_tensor_inputs = [
"latents",
"prompt_embeds",
"negative_prompt_embeds",
"prompt_embeds_2",
"negative_prompt_embeds_2",
]
def __init__(
self,
vae: AutoencoderKLMagvit,
text_encoder: BertModel,
tokenizer: BertTokenizer,
text_encoder_2: T5EncoderModel,
tokenizer_2: T5Tokenizer,
transformer: EasyAnimateTransformer3DModel,
scheduler: DDIMScheduler,
safety_checker: StableDiffusionSafetyChecker,
feature_extractor: CLIPImageProcessor,
requires_safety_checker: bool = True
):
super().__init__()
self.register_modules(
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
tokenizer_2=tokenizer_2,
transformer=transformer,
scheduler=scheduler,
safety_checker=safety_checker,
feature_extractor=feature_extractor,
text_encoder_2=text_encoder_2
)
if safety_checker is None and requires_safety_checker:
logger.warning(
f"You have disabled the safety checker for {self.__class__} by passing `safety_checker=None`. Ensure"
" that you abide to the conditions of the Stable Diffusion license and do not expose unfiltered"
" results in services or applications open to the public. Both the diffusers team and Hugging Face"
" strongly recommend to keep the safety filter enabled in all public facing circumstances, disabling"
" it only for use-cases that involve analyzing network behavior or auditing its results. For more"
" information, please have a look at https://github.com/huggingface/diffusers/pull/254 ."
)
if safety_checker is not None and feature_extractor is None:
raise ValueError(
"Make sure to define a feature extractor when loading {self.__class__} if you want to use the safety"
" checker. If you do not want to use the safety checker, you can pass `'safety_checker=None'` instead."
)
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
self.mask_processor = VaeImageProcessor(
vae_scale_factor=self.vae_scale_factor, do_normalize=False, do_binarize=True, do_convert_grayscale=True
)
self.register_to_config(requires_safety_checker=requires_safety_checker)
def enable_sequential_cpu_offload(self, *args, **kwargs):
super().enable_sequential_cpu_offload(*args, **kwargs)
if hasattr(self.transformer, "clip_projection") and self.transformer.clip_projection is not None:
import accelerate
accelerate.hooks.remove_hook_from_module(self.transformer.clip_projection, recurse=True)
self.transformer.clip_projection = self.transformer.clip_projection.to("cuda")
def encode_prompt(
self,
prompt: str,
device: torch.device,
dtype: torch.dtype,
num_images_per_prompt: int = 1,
do_classifier_free_guidance: bool = True,
negative_prompt: Optional[str] = None,
prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
prompt_attention_mask: Optional[torch.Tensor] = None,
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
max_sequence_length: Optional[int] = None,
text_encoder_index: int = 0,
actual_max_sequence_length: int = 256
):
r"""
Encodes the prompt into text encoder hidden states.
Args:
prompt (`str` or `List[str]`, *optional*):
prompt to be encoded
device: (`torch.device`):
torch device
dtype (`torch.dtype`):
torch dtype
num_images_per_prompt (`int`):
number of images that should be generated per prompt
do_classifier_free_guidance (`bool`):
whether to use classifier free guidance or not
negative_prompt (`str` or `List[str]`, *optional*):
The prompt or prompts not to guide the image generation. If not defined, one has to pass
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
less than `1`).
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
provided, text embeddings will be generated from `prompt` input argument.
negative_prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
argument.
prompt_attention_mask (`torch.Tensor`, *optional*):
Attention mask for the prompt. Required when `prompt_embeds` is passed directly.
negative_prompt_attention_mask (`torch.Tensor`, *optional*):
Attention mask for the negative prompt. Required when `negative_prompt_embeds` is passed directly.
max_sequence_length (`int`, *optional*): maximum sequence length to use for the prompt.
text_encoder_index (`int`, *optional*):
Index of the text encoder to use. `0` for clip and `1` for T5.
"""
tokenizers = [self.tokenizer, self.tokenizer_2]
text_encoders = [self.text_encoder, self.text_encoder_2]
tokenizer = tokenizers[text_encoder_index]
text_encoder = text_encoders[text_encoder_index]
if max_sequence_length is None:
if text_encoder_index == 0:
max_length = min(self.tokenizer.model_max_length, actual_max_sequence_length)
if text_encoder_index == 1:
max_length = min(self.tokenizer_2.model_max_length, actual_max_sequence_length)
else:
max_length = max_sequence_length
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
if prompt_embeds is None:
text_inputs = tokenizer(
prompt,
padding="max_length",
max_length=max_length,
truncation=True,
return_attention_mask=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
if text_input_ids.shape[-1] > actual_max_sequence_length:
reprompt = tokenizer.batch_decode(text_input_ids[:, :actual_max_sequence_length], skip_special_tokens=True)
text_inputs = tokenizer(
reprompt,
padding="max_length",
max_length=max_length,
truncation=True,
return_attention_mask=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
untruncated_ids = tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(
text_input_ids, untruncated_ids
):
_actual_max_sequence_length = min(tokenizer.model_max_length, actual_max_sequence_length)
removed_text = tokenizer.batch_decode(untruncated_ids[:, _actual_max_sequence_length - 1 : -1])
logger.warning(
"The following part of your input was truncated because CLIP can only handle sequences up to"
f" {_actual_max_sequence_length} tokens: {removed_text}"
)
prompt_attention_mask = text_inputs.attention_mask.to(device)
if self.transformer.config.enable_text_attention_mask:
prompt_embeds = text_encoder(
text_input_ids.to(device),
attention_mask=prompt_attention_mask,
)
else:
prompt_embeds = text_encoder(
text_input_ids.to(device)
)
prompt_embeds = prompt_embeds[0]
prompt_attention_mask = prompt_attention_mask.repeat(num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
bs_embed, seq_len, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1)
# get unconditional embeddings for classifier free guidance
if do_classifier_free_guidance and negative_prompt_embeds is None:
uncond_tokens: List[str]
if negative_prompt is None:
uncond_tokens = [""] * batch_size
elif prompt is not None and type(prompt) is not type(negative_prompt):
raise TypeError(
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
f" {type(prompt)}."
)
elif isinstance(negative_prompt, str):
uncond_tokens = [negative_prompt]
elif batch_size != len(negative_prompt):
raise ValueError(
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
" the batch size of `prompt`."
)
else:
uncond_tokens = negative_prompt
max_length = prompt_embeds.shape[1]
uncond_input = tokenizer(
uncond_tokens,
padding="max_length",
max_length=max_length,
truncation=True,
return_tensors="pt",
)
uncond_input_ids = uncond_input.input_ids
if uncond_input_ids.shape[-1] > actual_max_sequence_length:
reuncond_tokens = tokenizer.batch_decode(uncond_input_ids[:, :actual_max_sequence_length], skip_special_tokens=True)
uncond_input = tokenizer(
reuncond_tokens,
padding="max_length",
max_length=max_length,
truncation=True,
return_attention_mask=True,
return_tensors="pt",
)
uncond_input_ids = uncond_input.input_ids
negative_prompt_attention_mask = uncond_input.attention_mask.to(device)
if self.transformer.config.enable_text_attention_mask:
negative_prompt_embeds = text_encoder(
uncond_input.input_ids.to(device),
attention_mask=negative_prompt_attention_mask,
)
else:
negative_prompt_embeds = text_encoder(
uncond_input.input_ids.to(device)
)
negative_prompt_embeds = negative_prompt_embeds[0]
negative_prompt_attention_mask = negative_prompt_attention_mask.repeat(num_images_per_prompt, 1)
if do_classifier_free_guidance:
# duplicate unconditional embeddings for each generation per prompt, using mps friendly method
seq_len = negative_prompt_embeds.shape[1]
negative_prompt_embeds = negative_prompt_embeds.to(dtype=dtype, device=device)
negative_prompt_embeds = negative_prompt_embeds.repeat(1, num_images_per_prompt, 1)
negative_prompt_embeds = negative_prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1)
return prompt_embeds, negative_prompt_embeds, prompt_attention_mask, negative_prompt_attention_mask
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.run_safety_checker
def run_safety_checker(self, image, device, dtype):
if self.safety_checker is None:
has_nsfw_concept = None
else:
if torch.is_tensor(image):
feature_extractor_input = self.image_processor.postprocess(image, output_type="pil")
else:
feature_extractor_input = self.image_processor.numpy_to_pil(image)
safety_checker_input = self.feature_extractor(feature_extractor_input, return_tensors="pt").to(device)
image, has_nsfw_concept = self.safety_checker(
images=image, clip_input=safety_checker_input.pixel_values.to(dtype)
)
return image, has_nsfw_concept
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_extra_step_kwargs
def prepare_extra_step_kwargs(self, generator, eta):
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
# eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.
# eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502
# and should be between [0, 1]
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
extra_step_kwargs = {}
if accepts_eta:
extra_step_kwargs["eta"] = eta
# check if the scheduler accepts generator
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
if accepts_generator:
extra_step_kwargs["generator"] = generator
return extra_step_kwargs
def check_inputs(
self,
prompt,
height,
width,
negative_prompt=None,
prompt_embeds=None,
negative_prompt_embeds=None,
prompt_attention_mask=None,
negative_prompt_attention_mask=None,
prompt_embeds_2=None,
negative_prompt_embeds_2=None,
prompt_attention_mask_2=None,
negative_prompt_attention_mask_2=None,
callback_on_step_end_tensor_inputs=None,
):
if height % 8 != 0 or width % 8 != 0:
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
if prompt is not None and prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
" only forward one of the two."
)
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is None and prompt_embeds_2 is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds_2`. Cannot leave both `prompt` and `prompt_embeds_2` undefined."
)
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
if prompt_embeds is not None and prompt_attention_mask is None:
raise ValueError("Must provide `prompt_attention_mask` when specifying `prompt_embeds`.")
if prompt_embeds_2 is not None and prompt_attention_mask_2 is None:
raise ValueError("Must provide `prompt_attention_mask_2` when specifying `prompt_embeds_2`.")
if negative_prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
)
if negative_prompt_embeds is not None and negative_prompt_attention_mask is None:
raise ValueError("Must provide `negative_prompt_attention_mask` when specifying `negative_prompt_embeds`.")
if negative_prompt_embeds_2 is not None and negative_prompt_attention_mask_2 is None:
raise ValueError(
"Must provide `negative_prompt_attention_mask_2` when specifying `negative_prompt_embeds_2`."
)
if prompt_embeds is not None and negative_prompt_embeds is not None:
if prompt_embeds.shape != negative_prompt_embeds.shape:
raise ValueError(
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
f" {negative_prompt_embeds.shape}."
)
if prompt_embeds_2 is not None and negative_prompt_embeds_2 is not None:
if prompt_embeds_2.shape != negative_prompt_embeds_2.shape:
raise ValueError(
"`prompt_embeds_2` and `negative_prompt_embeds_2` must have the same shape when passed directly, but"
f" got: `prompt_embeds_2` {prompt_embeds_2.shape} != `negative_prompt_embeds_2`"
f" {negative_prompt_embeds_2.shape}."
)
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_latents
def prepare_latents(self, batch_size, num_channels_latents, video_length, height, width, dtype, device, generator, latents=None):
if self.vae.quant_conv is None or self.vae.quant_conv.weight.ndim==5:
if self.vae.cache_mag_vae:
mini_batch_encoder = self.vae.mini_batch_encoder
mini_batch_decoder = self.vae.mini_batch_decoder
shape = (batch_size, num_channels_latents, int((video_length - 1) // mini_batch_encoder * mini_batch_decoder + 1) if video_length != 1 else 1, height // self.vae_scale_factor, width // self.vae_scale_factor)
else:
mini_batch_encoder = self.vae.mini_batch_encoder
mini_batch_decoder = self.vae.mini_batch_decoder
shape = (batch_size, num_channels_latents, int(video_length // mini_batch_encoder * mini_batch_decoder) if video_length != 1 else 1, height // self.vae_scale_factor, width // self.vae_scale_factor)
else:
shape = (batch_size, num_channels_latents, video_length, height // self.vae_scale_factor, width // self.vae_scale_factor)
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
if latents is None:
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
else:
latents = latents.to(device)
# scale the initial noise by the standard deviation required by the scheduler
latents = latents * self.scheduler.init_noise_sigma
return latents
def prepare_control_latents(
self, mask, masked_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance
):
# resize the mask to latents shape as we concatenate the mask to the latents
# we do that before converting to dtype to avoid breaking in case we're using cpu_offload
# and half precision
if mask is not None:
mask = mask.to(device=device, dtype=dtype)
bs = 1
new_mask = []
for i in range(0, mask.shape[0], bs):
mask_bs = mask[i : i + bs]
mask_bs = self.vae.encode(mask_bs)[0]
mask_bs = mask_bs.mode()
new_mask.append(mask_bs)
mask = torch.cat(new_mask, dim = 0)
mask = mask * self.vae.config.scaling_factor
if masked_image is not None:
masked_image = masked_image.to(device=device, dtype=dtype)
bs = 1
new_mask_pixel_values = []
for i in range(0, masked_image.shape[0], bs):
mask_pixel_values_bs = masked_image[i : i + bs]
mask_pixel_values_bs = self.vae.encode(mask_pixel_values_bs)[0]
mask_pixel_values_bs = mask_pixel_values_bs.mode()
new_mask_pixel_values.append(mask_pixel_values_bs)
masked_image_latents = torch.cat(new_mask_pixel_values, dim = 0)
masked_image_latents = masked_image_latents * self.vae.config.scaling_factor
else:
masked_image_latents = None
return mask, masked_image_latents
def smooth_output(self, video, mini_batch_encoder, mini_batch_decoder):
if video.size()[2] <= mini_batch_encoder:
return video
prefix_index_before = mini_batch_encoder // 2
prefix_index_after = mini_batch_encoder - prefix_index_before
pixel_values = video[:, :, prefix_index_before:-prefix_index_after]
# Encode middle videos
latents = self.vae.encode(pixel_values)[0]
latents = latents.mode()
# Decode middle videos
middle_video = self.vae.decode(latents)[0]
video[:, :, prefix_index_before:-prefix_index_after] = (video[:, :, prefix_index_before:-prefix_index_after] + middle_video) / 2
return video
def decode_latents(self, latents):
video_length = latents.shape[2]
latents = 1 / self.vae.config.scaling_factor * latents
if self.vae.quant_conv is None or self.vae.quant_conv.weight.ndim==5:
mini_batch_encoder = self.vae.mini_batch_encoder
mini_batch_decoder = self.vae.mini_batch_decoder
video = self.vae.decode(latents)[0]
video = video.clamp(-1, 1)
if not self.vae.cache_compression_vae and not self.vae.cache_mag_vae:
video = self.smooth_output(video, mini_batch_encoder, mini_batch_decoder).cpu().clamp(-1, 1)
else:
latents = rearrange(latents, "b c f h w -> (b f) c h w")
video = []
for frame_idx in tqdm(range(latents.shape[0])):
video.append(self.vae.decode(latents[frame_idx:frame_idx+1]).sample)
video = torch.cat(video)
video = rearrange(video, "(b f) c h w -> b c f h w", f=video_length)
video = (video / 2 + 0.5).clamp(0, 1)
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16
video = video.cpu().float().numpy()
return video
@property
def guidance_scale(self):
return self._guidance_scale
@property
def guidance_rescale(self):
return self._guidance_rescale
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
# corresponds to doing no classifier free guidance.
@property
def do_classifier_free_guidance(self):
return self._guidance_scale > 1
@property
def num_timesteps(self):
return self._num_timesteps
@property
def interrupt(self):
return self._interrupt
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(
self,
prompt: Union[str, List[str]] = None,
video_length: Optional[int] = None,
height: Optional[int] = None,
width: Optional[int] = None,
control_video: Union[torch.FloatTensor] = None,
num_inference_steps: Optional[int] = 50,
guidance_scale: Optional[float] = 5.0,
negative_prompt: Optional[Union[str, List[str]]] = None,
num_images_per_prompt: Optional[int] = 1,
eta: Optional[float] = 0.0,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
prompt_embeds: Optional[torch.Tensor] = None,
prompt_embeds_2: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds_2: Optional[torch.Tensor] = None,
prompt_attention_mask: Optional[torch.Tensor] = None,
prompt_attention_mask_2: Optional[torch.Tensor] = None,
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
negative_prompt_attention_mask_2: Optional[torch.Tensor] = None,
output_type: Optional[str] = "latent",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
guidance_rescale: float = 0.0,
original_size: Optional[Tuple[int, int]] = (1024, 1024),
target_size: Optional[Tuple[int, int]] = None,
crops_coords_top_left: Tuple[int, int] = (0, 0),
comfyui_progressbar: bool = False,
):
r"""
Generates images or video using the EasyAnimate pipeline based on the provided prompts.
Examples:
prompt (`str` or `List[str]`, *optional*):
Text prompts to guide the image or video generation. If not provided, use `prompt_embeds` instead.
video_length (`int`, *optional*):
Length of the generated video (in frames).
height (`int`, *optional*):
Height of the generated image in pixels.
width (`int`, *optional*):
Width of the generated image in pixels.
num_inference_steps (`int`, *optional*, defaults to 50):
Number of denoising steps during generation. More steps generally yield higher quality images but slow down inference.
guidance_scale (`float`, *optional*, defaults to 5.0):
Encourages the model to align outputs with prompts. A higher value may decrease image quality.
negative_prompt (`str` or `List[str]`, *optional*):
Prompts indicating what to exclude in generation. If not specified, use `negative_prompt_embeds`.
num_images_per_prompt (`int`, *optional*, defaults to 1):
Number of images to generate for each prompt.
eta (`float`, *optional*, defaults to 0.0):
Applies to DDIM scheduling. Controlled by the eta parameter from the related literature.
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
A generator to ensure reproducibility in image generation.
latents (`torch.Tensor`, *optional*):
Predefined latent tensors to condition generation.
prompt_embeds (`torch.Tensor`, *optional*):
Text embeddings for the prompts. Overrides prompt string inputs for more flexibility.
prompt_embeds_2 (`torch.Tensor`, *optional*):
Secondary text embeddings to supplement or replace the initial prompt embeddings.
negative_prompt_embeds (`torch.Tensor`, *optional*):
Embeddings for negative prompts. Overrides string inputs if defined.
negative_prompt_embeds_2 (`torch.Tensor`, *optional*):
Secondary embeddings for negative prompts, similar to `negative_prompt_embeds`.
prompt_attention_mask (`torch.Tensor`, *optional*):
Attention mask for the primary prompt embeddings.
prompt_attention_mask_2 (`torch.Tensor`, *optional*):
Attention mask for the secondary prompt embeddings.
negative_prompt_attention_mask (`torch.Tensor`, *optional*):
Attention mask for negative prompt embeddings.
negative_prompt_attention_mask_2 (`torch.Tensor`, *optional*):
Attention mask for secondary negative prompt embeddings.
output_type (`str`, *optional*, defaults to "latent"):
Format of the generated output, either as a PIL image or as a NumPy array.
return_dict (`bool`, *optional*, defaults to `True`):
If `True`, returns a structured output. Otherwise returns a simple tuple.
callback_on_step_end (`Callable`, *optional*):
Functions called at the end of each denoising step.
callback_on_step_end_tensor_inputs (`List[str]`, *optional*):
Tensor names to be included in callback function calls.
guidance_rescale (`float`, *optional*, defaults to 0.0):
Adjusts noise levels based on guidance scale.
original_size (`Tuple[int, int]`, *optional*, defaults to `(1024, 1024)`):
Original dimensions of the output.
target_size (`Tuple[int, int]`, *optional*):
Desired output dimensions for calculations.
crops_coords_top_left (`Tuple[int, int]`, *optional*, defaults to `(0, 0)`):
Coordinates for cropping.
Returns:
[`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] or `tuple`:
If `return_dict` is `True`, [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] is returned,
otherwise a `tuple` is returned where the first element is a list with the generated images and the
second element is a list of `bool`s indicating whether the corresponding generated image contains
"not-safe-for-work" (nsfw) content.
"""
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
# 0. default height and width
height = int((height // 16) * 16)
width = int((width // 16) * 16)
# 1. Check inputs. Raise error if not correct
self.check_inputs(
prompt,
height,
width,
negative_prompt,
prompt_embeds,
negative_prompt_embeds,
prompt_attention_mask,
negative_prompt_attention_mask,
prompt_embeds_2,
negative_prompt_embeds_2,
prompt_attention_mask_2,
negative_prompt_attention_mask_2,
callback_on_step_end_tensor_inputs,
)
self._guidance_scale = guidance_scale
self._guidance_rescale = guidance_rescale
self._interrupt = False
# 2. Define call parameters
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
device = self._execution_device
if self.text_encoder is not None:
dtype = self.text_encoder.dtype
elif self.text_encoder_2 is not None:
dtype = self.text_encoder_2.dtype
else:
dtype = self.transformer.dtype
# 3. Encode input prompt
(
prompt_embeds,
negative_prompt_embeds,
prompt_attention_mask,
negative_prompt_attention_mask,
) = self.encode_prompt(
prompt=prompt,
device=device,
dtype=dtype,
num_images_per_prompt=num_images_per_prompt,
do_classifier_free_guidance=self.do_classifier_free_guidance,
negative_prompt=negative_prompt,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
prompt_attention_mask=prompt_attention_mask,
negative_prompt_attention_mask=negative_prompt_attention_mask,
text_encoder_index=0,
)
(
prompt_embeds_2,
negative_prompt_embeds_2,
prompt_attention_mask_2,
negative_prompt_attention_mask_2,
) = self.encode_prompt(
prompt=prompt,
device=device,
dtype=dtype,
num_images_per_prompt=num_images_per_prompt,
do_classifier_free_guidance=self.do_classifier_free_guidance,
negative_prompt=negative_prompt,
prompt_embeds=prompt_embeds_2,
negative_prompt_embeds=negative_prompt_embeds_2,
prompt_attention_mask=prompt_attention_mask_2,
negative_prompt_attention_mask=negative_prompt_attention_mask_2,
text_encoder_index=1,
)
# 4. Prepare timesteps
self.scheduler.set_timesteps(num_inference_steps, device=device)
timesteps = self.scheduler.timesteps
if comfyui_progressbar:
from comfy.utils import ProgressBar
pbar = ProgressBar(num_inference_steps + 2)
# 5. Prepare latent variables
num_channels_latents = self.vae.config.latent_channels
latents = self.prepare_latents(
batch_size * num_images_per_prompt,
num_channels_latents,
video_length,
height,
width,
dtype,
device,
generator,
latents,
)
if comfyui_progressbar:
pbar.update(1)
if control_video is not None:
video_length = control_video.shape[2]
control_video = self.image_processor.preprocess(rearrange(control_video, "b c f h w -> (b f) c h w"), height=height, width=width)
control_video = control_video.to(dtype=torch.float32)
control_video = rearrange(control_video, "(b f) c h w -> b c f h w", f=video_length)
else:
control_video = None
control_video_latents = self.prepare_control_latents(
None,
control_video,
batch_size,
height,
width,
dtype,
device,
generator,
self.do_classifier_free_guidance
)[1]
control_latents = (
torch.cat([control_video_latents] * 2) if self.do_classifier_free_guidance else control_video_latents
)
if comfyui_progressbar:
pbar.update(1)
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
# 7 create image_rotary_emb, style embedding & time ids
grid_height = height // 8 // self.transformer.config.patch_size
grid_width = width // 8 // self.transformer.config.patch_size
if self.transformer.config.get("time_position_encoding_type", "2d_rope") == "3d_rope":
base_size_width = 720 // 8 // self.transformer.config.patch_size
base_size_height = 480 // 8 // self.transformer.config.patch_size
grid_crops_coords = get_resize_crop_region_for_grid(
(grid_height, grid_width), base_size_width, base_size_height
)
image_rotary_emb = get_3d_rotary_pos_embed(
self.transformer.config.attention_head_dim, grid_crops_coords, grid_size=(grid_height, grid_width),
temporal_size=latents.size(2), use_real=True,
)
else:
base_size = 512 // 8 // self.transformer.config.patch_size
grid_crops_coords = get_resize_crop_region_for_grid(
(grid_height, grid_width), base_size, base_size
)
image_rotary_emb = get_2d_rotary_pos_embed(
self.transformer.config.attention_head_dim, grid_crops_coords, (grid_height, grid_width)
)
# Get other hunyuan params
style = torch.tensor([0], device=device)
target_size = target_size or (height, width)
add_time_ids = list(original_size + target_size + crops_coords_top_left)
add_time_ids = torch.tensor([add_time_ids], dtype=dtype)
if self.do_classifier_free_guidance:
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds])
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask])
prompt_embeds_2 = torch.cat([negative_prompt_embeds_2, prompt_embeds_2])
prompt_attention_mask_2 = torch.cat([negative_prompt_attention_mask_2, prompt_attention_mask_2])
add_time_ids = torch.cat([add_time_ids] * 2, dim=0)
style = torch.cat([style] * 2, dim=0)
# To latents.device
prompt_embeds = prompt_embeds.to(device=device)
prompt_attention_mask = prompt_attention_mask.to(device=device)
prompt_embeds_2 = prompt_embeds_2.to(device=device)
prompt_attention_mask_2 = prompt_attention_mask_2.to(device=device)
add_time_ids = add_time_ids.to(dtype=dtype, device=device).repeat(
batch_size * num_images_per_prompt, 1
)
style = style.to(device=device).repeat(batch_size * num_images_per_prompt)
# 8. Denoising loop
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
self._num_timesteps = len(timesteps)
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if self.interrupt:
continue
# expand the latents if we are doing classifier free guidance
latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
# expand scalar t to 1-D tensor to match the 1st dim of latent_model_input
t_expand = torch.tensor([t] * latent_model_input.shape[0], device=device).to(
dtype=latent_model_input.dtype
)
# predict the noise residual
noise_pred = self.transformer(
latent_model_input,
t_expand,
encoder_hidden_states=prompt_embeds,
text_embedding_mask=prompt_attention_mask,
encoder_hidden_states_t5=prompt_embeds_2,
text_embedding_mask_t5=prompt_attention_mask_2,
image_meta_size=add_time_ids,
style=style,
image_rotary_emb=image_rotary_emb,
return_dict=False,
control_latents=control_latents,
)[0]
if noise_pred.size()[1] != self.vae.config.latent_channels:
noise_pred, _ = noise_pred.chunk(2, dim=1)
# perform guidance
if self.do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
if self.do_classifier_free_guidance and guidance_rescale > 0.0:
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=guidance_rescale)
# compute the previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
if callback_on_step_end is not None:
callback_kwargs = {}
for k in callback_on_step_end_tensor_inputs:
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
prompt_embeds_2 = callback_outputs.pop("prompt_embeds_2", prompt_embeds_2)
negative_prompt_embeds_2 = callback_outputs.pop(
"negative_prompt_embeds_2", negative_prompt_embeds_2
)
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if XLA_AVAILABLE:
xm.mark_step()
if comfyui_progressbar:
pbar.update(1)
# Post-processing
video = self.decode_latents(latents)
# Convert to tensor
if output_type == "latent":
video = torch.from_numpy(video)
# Offload all models
self.maybe_free_model_hooks()
if not return_dict:
return video
return EasyAnimatePipelineOutput(videos=video)
File diff suppressed because it is too large Load Diff
+1
View File
@@ -0,0 +1 @@
This folder is modified from the official [MPS](https://github.com/Kwai-Kolors/MPS/tree/main) repository.
@@ -0,0 +1,7 @@
from dataclasses import dataclass
@dataclass
class BaseModelConfig:
pass
@@ -0,0 +1,154 @@
from dataclasses import dataclass
from transformers import CLIPModel as HFCLIPModel
from transformers import AutoTokenizer
from torch import nn, einsum
# Modified: import
# from trainer.models.base_model import BaseModelConfig
from .base_model import BaseModelConfig
from transformers import CLIPConfig
from typing import Any, Optional, Tuple, Union
import torch
# Modified: import
# from trainer.models.cross_modeling import Cross_model
from .cross_modeling import Cross_model
import gc
class XCLIPModel(HFCLIPModel):
def __init__(self, config: CLIPConfig):
super().__init__(config)
def get_text_features(
self,
input_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> torch.FloatTensor:
# Use CLIP model's config for some fields (if specified) instead of those of vision & text components.
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
text_outputs = self.text_model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
# pooled_output = text_outputs[1]
# text_features = self.text_projection(pooled_output)
last_hidden_state = text_outputs[0]
text_features = self.text_projection(last_hidden_state)
pooled_output = text_outputs[1]
text_features_EOS = self.text_projection(pooled_output)
# del last_hidden_state, text_outputs
# gc.collect()
return text_features, text_features_EOS
def get_image_features(
self,
pixel_values: Optional[torch.FloatTensor] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> torch.FloatTensor:
# Use CLIP model's config for some fields (if specified) instead of those of vision & text components.
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
vision_outputs = self.vision_model(
pixel_values=pixel_values,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
# pooled_output = vision_outputs[1] # pooled_output
# image_features = self.visual_projection(pooled_output)
last_hidden_state = vision_outputs[0]
image_features = self.visual_projection(last_hidden_state)
return image_features
@dataclass
class ClipModelConfig(BaseModelConfig):
_target_: str = "trainer.models.clip_model.CLIPModel"
pretrained_model_name_or_path: str ="openai/clip-vit-base-patch32"
class CLIPModel(nn.Module):
def __init__(self, config):
super().__init__()
# Modified: We convert the original ckpt (contains the entire model) to a `state_dict`.
# self.model = XCLIPModel.from_pretrained(ckpt)
self.model = XCLIPModel(config)
self.cross_model = Cross_model(dim=1024, layer_num=4, heads=16)
def get_text_features(self, *args, **kwargs):
return self.model.get_text_features(*args, **kwargs)
def get_image_features(self, *args, **kwargs):
return self.model.get_image_features(*args, **kwargs)
def forward(self, text_inputs=None, image_inputs=None, condition_inputs=None):
outputs = ()
text_f, text_EOS = self.model.get_text_features(text_inputs) # B*77*1024
outputs += text_EOS,
image_f = self.model.get_image_features(image_inputs.half()) # 2B*257*1024
# [B, 77, 1024]
condition_f, _ = self.model.get_text_features(condition_inputs) # B*5*1024
sim_text_condition = einsum('b i d, b j d -> b j i', text_f, condition_f)
sim_text_condition = torch.max(sim_text_condition, dim=1, keepdim=True)[0]
sim_text_condition = sim_text_condition / sim_text_condition.max()
mask = torch.where(sim_text_condition > 0.01, 0, float('-inf')) # B*1*77
# Modified: Support both torch.float16 and torch.bfloat16
# mask = mask.repeat(1,image_f.shape[1],1) # B*257*77
model_dtype = next(self.cross_model.parameters()).dtype
mask = mask.repeat(1,image_f.shape[1],1).to(model_dtype) # B*257*77
# bc = int(image_f.shape[0]/2)
# Modified: The original input consists of a (batch of) text and two (batches of) images,
# primarily used to compute which (batch of) image is more consistent with the text.
# The modified input consists of a (batch of) text and a (batch of) images.
# sim0 = self.cross_model(image_f[:bc,:,:], text_f,mask.half())
# sim1 = self.cross_model(image_f[bc:,:,:], text_f,mask.half())
# outputs += sim0[:,0,:],
# outputs += sim1[:,0,:],
sim = self.cross_model(image_f, text_f,mask)
outputs += sim[:,0,:],
return outputs
@property
def logit_scale(self):
return self.model.logit_scale
def save(self, path):
self.model.save_pretrained(path)
@@ -0,0 +1,291 @@
import torch
from torch import einsum, nn
import torch.nn.functional as F
from einops import rearrange, repeat
# helper functions
def exists(val):
return val is not None
def default(val, d):
return val if exists(val) else d
# normalization
# they use layernorm without bias, something that pytorch does not offer
class LayerNorm(nn.Module):
def __init__(self, dim):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.register_buffer("bias", torch.zeros(dim))
def forward(self, x):
return F.layer_norm(x, x.shape[-1:], self.weight, self.bias)
# residual
class Residual(nn.Module):
def __init__(self, fn):
super().__init__()
self.fn = fn
def forward(self, x, *args, **kwargs):
return self.fn(x, *args, **kwargs) + x
# rotary positional embedding
# https://arxiv.org/abs/2104.09864
class RotaryEmbedding(nn.Module):
def __init__(self, dim):
super().__init__()
inv_freq = 1.0 / (10000 ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer("inv_freq", inv_freq)
def forward(self, max_seq_len, *, device):
seq = torch.arange(max_seq_len, device=device, dtype=self.inv_freq.dtype)
freqs = einsum("i , j -> i j", seq, self.inv_freq)
return torch.cat((freqs, freqs), dim=-1)
def rotate_half(x):
x = rearrange(x, "... (j d) -> ... j d", j=2)
x1, x2 = x.unbind(dim=-2)
return torch.cat((-x2, x1), dim=-1)
def apply_rotary_pos_emb(pos, t):
return (t * pos.cos()) + (rotate_half(t) * pos.sin())
# classic Noam Shazeer paper, except here they use SwiGLU instead of the more popular GEGLU for gating the feedforward
# https://arxiv.org/abs/2002.05202
class SwiGLU(nn.Module):
def forward(self, x):
x, gate = x.chunk(2, dim=-1)
return F.silu(gate) * x
# parallel attention and feedforward with residual
# discovered by Wang et al + EleutherAI from GPT-J fame
class ParallelTransformerBlock(nn.Module):
def __init__(self, dim, dim_head=64, heads=8, ff_mult=4):
super().__init__()
self.norm = LayerNorm(dim)
attn_inner_dim = dim_head * heads
ff_inner_dim = dim * ff_mult
self.fused_dims = (attn_inner_dim, dim_head, dim_head, (ff_inner_dim * 2))
self.heads = heads
self.scale = dim_head**-0.5
self.rotary_emb = RotaryEmbedding(dim_head)
self.fused_attn_ff_proj = nn.Linear(dim, sum(self.fused_dims), bias=False)
self.attn_out = nn.Linear(attn_inner_dim, dim, bias=False)
self.ff_out = nn.Sequential(
SwiGLU(),
nn.Linear(ff_inner_dim, dim, bias=False)
)
self.register_buffer("pos_emb", None, persistent=False)
def get_rotary_embedding(self, n, device):
if self.pos_emb is not None and self.pos_emb.shape[-2] >= n:
return self.pos_emb[:n]
pos_emb = self.rotary_emb(n, device=device)
self.register_buffer("pos_emb", pos_emb, persistent=False)
return pos_emb
def forward(self, x, attn_mask=None):
"""
einstein notation
b - batch
h - heads
n, i, j - sequence length (base sequence length, source, target)
d - feature dimension
"""
n, device, h = x.shape[1], x.device, self.heads
# pre layernorm
x = self.norm(x)
# attention queries, keys, values, and feedforward inner
q, k, v, ff = self.fused_attn_ff_proj(x).split(self.fused_dims, dim=-1)
# split heads
# they use multi-query single-key-value attention, yet another Noam Shazeer paper
# they found no performance loss past a certain scale, and more efficient decoding obviously
# https://arxiv.org/abs/1911.02150
q = rearrange(q, "b n (h d) -> b h n d", h=h)
# rotary embeddings
positions = self.get_rotary_embedding(n, device)
q, k = map(lambda t: apply_rotary_pos_emb(positions, t), (q, k))
# scale
q = q * self.scale
# similarity
sim = einsum("b h i d, b j d -> b h i j", q, k)
# extra attention mask - for masking out attention from text CLS token to padding
if exists(attn_mask):
attn_mask = rearrange(attn_mask, 'b i j -> b 1 i j')
sim = sim.masked_fill(~attn_mask, -torch.finfo(sim.dtype).max)
# attention
sim = sim - sim.amax(dim=-1, keepdim=True).detach()
attn = sim.softmax(dim=-1)
# aggregate values
out = einsum("b h i j, b j d -> b h i d", attn, v)
# merge heads
out = rearrange(out, "b h n d -> b n (h d)")
return self.attn_out(out) + self.ff_out(ff)
# cross attention - using multi-query + one-headed key / values as in PaLM w/ optional parallel feedforward
class CrossAttention(nn.Module):
def __init__(
self,
dim,
*,
context_dim=None,
dim_head=64,
heads=12,
parallel_ff=False,
ff_mult=4,
norm_context=False
):
super().__init__()
self.heads = heads
self.scale = dim_head ** -0.5
inner_dim = heads * dim_head
context_dim = default(context_dim, dim)
self.norm = LayerNorm(dim)
self.context_norm = LayerNorm(context_dim) if norm_context else nn.Identity()
self.to_q = nn.Linear(dim, inner_dim, bias=False)
self.to_kv = nn.Linear(context_dim, dim_head * 2, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
# whether to have parallel feedforward
ff_inner_dim = ff_mult * dim
self.ff = nn.Sequential(
nn.Linear(dim, ff_inner_dim * 2, bias=False),
SwiGLU(),
nn.Linear(ff_inner_dim, dim, bias=False)
) if parallel_ff else None
def forward(self, x, context, mask):
"""
einstein notation
b - batch
h - heads
n, i, j - sequence length (base sequence length, source, target)
d - feature dimension
"""
# pre-layernorm, for queries and context
x = self.norm(x)
context = self.context_norm(context)
# get queries
q = self.to_q(x)
q = rearrange(q, 'b n (h d) -> b h n d', h = self.heads)
# scale
q = q * self.scale
# get key / values
k, v = self.to_kv(context).chunk(2, dim=-1)
# query / key similarity
sim = einsum('b h i d, b j d -> b h i j', q, k)
# attention
mask = mask.unsqueeze(1).repeat(1,self.heads,1,1)
sim = sim + mask # context mask
sim = sim - sim.amax(dim=-1, keepdim=True)
attn = sim.softmax(dim=-1)
# aggregate
out = einsum('b h i j, b j d -> b h i d', attn, v)
# merge and combine heads
out = rearrange(out, 'b h n d -> b n (h d)')
out = self.to_out(out)
# add parallel feedforward (for multimodal layers)
if exists(self.ff):
out = out + self.ff(x)
return out
class Cross_model(nn.Module):
def __init__(
self,
dim=512,
layer_num=4,
dim_head=64,
heads=8,
ff_mult=4
):
super().__init__()
self.layers = nn.ModuleList([])
for ind in range(layer_num):
self.layers.append(nn.ModuleList([
Residual(CrossAttention(dim=dim, dim_head=dim_head, heads=heads, parallel_ff=True, ff_mult=ff_mult)),
Residual(ParallelTransformerBlock(dim=dim, dim_head=dim_head, heads=heads, ff_mult=ff_mult))
]))
def forward(
self,
query_tokens,
context_tokens,
mask
):
for cross_attn, self_attn_ff in self.layers:
query_tokens = cross_attn(query_tokens, context_tokens,mask)
query_tokens = self_attn_ff(query_tokens)
return query_tokens
@@ -0,0 +1,13 @@
from .siglip_v2_5 import (
AestheticPredictorV2_5Head,
AestheticPredictorV2_5Model,
AestheticPredictorV2_5Processor,
convert_v2_5_from_siglip,
)
__all__ = [
"AestheticPredictorV2_5Head",
"AestheticPredictorV2_5Model",
"AestheticPredictorV2_5Processor",
"convert_v2_5_from_siglip",
]
@@ -0,0 +1,133 @@
# Borrowed from https://github.com/discus0434/aesthetic-predictor-v2-5/blob/3125a9e/src/aesthetic_predictor_v2_5/siglip_v2_5.py
import os
from collections import OrderedDict
from os import PathLike
from typing import Final
import torch
import torch.nn as nn
import torchvision.transforms as transforms
from transformers import (
SiglipImageProcessor,
SiglipVisionConfig,
SiglipVisionModel,
logging,
)
from transformers.image_processing_utils import BatchFeature
from transformers.modeling_outputs import ImageClassifierOutputWithNoAttention
logging.set_verbosity_error()
URL: Final[str] = (
"https://github.com/discus0434/aesthetic-predictor-v2-5/raw/main/models/aesthetic_predictor_v2_5.pth"
)
class AestheticPredictorV2_5Head(nn.Module):
def __init__(self, config: SiglipVisionConfig) -> None:
super().__init__()
self.scoring_head = nn.Sequential(
nn.Linear(config.hidden_size, 1024),
nn.Dropout(0.5),
nn.Linear(1024, 128),
nn.Dropout(0.5),
nn.Linear(128, 64),
nn.Dropout(0.5),
nn.Linear(64, 16),
nn.Dropout(0.2),
nn.Linear(16, 1),
)
def forward(self, image_embeds: torch.Tensor) -> torch.Tensor:
return self.scoring_head(image_embeds)
class AestheticPredictorV2_5Model(SiglipVisionModel):
PATCH_SIZE = 14
def __init__(self, config: SiglipVisionConfig, *args, **kwargs) -> None:
super().__init__(config, *args, **kwargs)
self.layers = AestheticPredictorV2_5Head(config)
self.post_init()
self.transforms = transforms.Compose([
transforms.Resize((384, 384)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
])
def forward(
self,
pixel_values: torch.FloatTensor | None = None,
labels: torch.Tensor | None = None,
return_dict: bool | None = None,
) -> tuple | ImageClassifierOutputWithNoAttention:
return_dict = (
return_dict if return_dict is not None else self.config.use_return_dict
)
outputs = super().forward(
pixel_values=pixel_values,
return_dict=return_dict,
)
image_embeds = outputs.pooler_output
image_embeds_norm = image_embeds / image_embeds.norm(dim=-1, keepdim=True)
prediction = self.layers(image_embeds_norm)
loss = None
if labels is not None:
loss_fct = nn.MSELoss()
loss = loss_fct()
if not return_dict:
return (loss, prediction, image_embeds)
return ImageClassifierOutputWithNoAttention(
loss=loss,
logits=prediction,
hidden_states=image_embeds,
)
class AestheticPredictorV2_5Processor(SiglipImageProcessor):
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
def __call__(self, *args, **kwargs) -> BatchFeature:
return super().__call__(*args, **kwargs)
@classmethod
def from_pretrained(
self,
pretrained_model_name_or_path: str
| PathLike = "google/siglip-so400m-patch14-384",
*args,
**kwargs,
) -> "AestheticPredictorV2_5Processor":
return super().from_pretrained(pretrained_model_name_or_path, *args, **kwargs)
def convert_v2_5_from_siglip(
predictor_name_or_path: str | PathLike | None = None,
encoder_model_name: str = "google/siglip-so400m-patch14-384",
*args,
**kwargs,
) -> tuple[AestheticPredictorV2_5Model, AestheticPredictorV2_5Processor]:
model = AestheticPredictorV2_5Model.from_pretrained(
encoder_model_name, *args, **kwargs
)
processor = AestheticPredictorV2_5Processor.from_pretrained(
encoder_model_name, *args, **kwargs
)
if predictor_name_or_path is None or not os.path.exists(predictor_name_or_path):
state_dict = torch.hub.load_state_dict_from_url(URL, map_location="cpu")
else:
state_dict = torch.load(predictor_name_or_path, map_location="cpu")
assert isinstance(state_dict, OrderedDict)
model.layers.load_state_dict(state_dict)
model.eval()
return model, processor
@@ -0,0 +1,49 @@
import os
import torch
import torch.nn as nn
from transformers import CLIPModel
from torchvision.datasets.utils import download_url
URL = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Third_Party/sac%2Blogos%2Bava1-l14-linearMSE.pth"
FILENAME = "sac+logos+ava1-l14-linearMSE.pth"
MD5 = "b1047fd767a00134b8fd6529bf19521a"
class MLP(nn.Module):
def __init__(self):
super().__init__()
self.layers = nn.Sequential(
nn.Linear(768, 1024),
nn.Dropout(0.2),
nn.Linear(1024, 128),
nn.Dropout(0.2),
nn.Linear(128, 64),
nn.Dropout(0.1),
nn.Linear(64, 16),
nn.Linear(16, 1),
)
def forward(self, embed):
return self.layers(embed)
class ImprovedAestheticPredictor(nn.Module):
def __init__(self, encoder_path="openai/clip-vit-large-patch14", predictor_path=None):
super().__init__()
self.encoder = CLIPModel.from_pretrained(encoder_path)
self.predictor = MLP()
if predictor_path is None or not os.path.exists(predictor_path):
download_url(URL, torch.hub.get_dir(), FILENAME, md5=MD5)
predictor_path = os.path.join(torch.hub.get_dir(), FILENAME)
state_dict = torch.load(predictor_path, map_location="cpu")
self.predictor.load_state_dict(state_dict)
self.eval()
def forward(self, pixel_values):
embed = self.encoder.get_image_features(pixel_values=pixel_values)
embed = embed / torch.linalg.vector_norm(embed, dim=-1, keepdim=True)
return self.predictor(embed).squeeze(1)
+385
View File
@@ -0,0 +1,385 @@
import os
from abc import ABC, abstractmethod
import torch
import torchvision.transforms as transforms
from einops import rearrange
from torchvision.datasets.utils import download_url
from typing import Optional, Tuple
# All reward models.
__all__ = ["AestheticReward", "HPSReward", "PickScoreReward", "MPSReward"]
class BaseReward(ABC):
"""An base class for reward models. A custom Reward class must implement two functions below.
"""
def __init__(self):
"""Define your reward model and image transformations (optional) here.
"""
pass
@abstractmethod
def __call__(self, batch_frames: torch.Tensor, batch_prompt: Optional[list[str]]=None) -> Tuple[torch.Tensor, torch.Tensor]:
"""Given batch frames with shape `[B, C, T, H, W]` extracted from a list of videos and a list of prompts
(optional) correspondingly, return the loss and reward computed by your reward model (reduction by mean).
"""
pass
class AestheticReward(BaseReward):
"""Aesthetic Predictor [V2](https://github.com/christophschuhmann/improved-aesthetic-predictor)
and [V2.5](https://github.com/discus0434/aesthetic-predictor-v2-5) reward model.
"""
def __init__(
self,
encoder_path="openai/clip-vit-large-patch14",
predictor_path=None,
version="v2",
device="cpu",
dtype=torch.float16,
max_reward=10,
loss_scale=0.1,
):
from .improved_aesthetic_predictor import ImprovedAestheticPredictor
from ..video_caption.utils.siglip_v2_5 import convert_v2_5_from_siglip
self.encoder_path = encoder_path
self.predictor_path = predictor_path
self.version = version
self.device = device
self.dtype = dtype
self.max_reward = max_reward
self.loss_scale = loss_scale
if self.version != "v2" and self.version != "v2.5":
raise ValueError("Only v2 and v2.5 are supported.")
if self.version == "v2":
assert "clip-vit-large-patch14" in encoder_path.lower()
self.model = ImprovedAestheticPredictor(encoder_path=self.encoder_path, predictor_path=self.predictor_path)
# https://huggingface.co/openai/clip-vit-large-patch14/blob/main/preprocessor_config.json
# TODO: [transforms.Resize(224), transforms.CenterCrop(224)] for any aspect ratio.
self.transform = transforms.Compose([
transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BICUBIC),
transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]),
])
elif self.version == "v2.5":
assert "siglip-so400m-patch14-384" in encoder_path.lower()
self.model, _ = convert_v2_5_from_siglip(encoder_model_name=self.encoder_path)
# https://huggingface.co/google/siglip-so400m-patch14-384/blob/main/preprocessor_config.json
self.transform = transforms.Compose([
transforms.Resize((384, 384), interpolation=transforms.InterpolationMode.BICUBIC),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
])
self.model.to(device=self.device, dtype=self.dtype)
self.model.requires_grad_(False)
def __call__(self, batch_frames: torch.Tensor, batch_prompt: Optional[list[str]]=None) -> Tuple[torch.Tensor, torch.Tensor]:
batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w")
batch_loss, batch_reward = 0, 0
for frames in batch_frames:
pixel_values = torch.stack([self.transform(frame) for frame in frames])
pixel_values = pixel_values.to(self.device, dtype=self.dtype)
if self.version == "v2":
reward = self.model(pixel_values)
elif self.version == "v2.5":
reward = self.model(pixel_values).logits.squeeze()
# Convert reward to loss in [0, 1].
if self.max_reward is None:
loss = (-1 * reward) * self.loss_scale
else:
loss = abs(reward - self.max_reward) * self.loss_scale
batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean()
return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0]
class HPSReward(BaseReward):
"""[HPS](https://github.com/tgxs002/HPSv2) v2 and v2.1 reward model.
"""
def __init__(
self,
model_path=None,
version="v2.0",
device="cpu",
dtype=torch.float16,
max_reward=1,
loss_scale=1,
):
from hpsv2.src.open_clip import create_model_and_transforms, get_tokenizer
self.model_path = model_path
self.version = version
self.device = device
self.dtype = dtype
self.max_reward = max_reward
self.loss_scale = loss_scale
self.model, _, _ = create_model_and_transforms(
"ViT-H-14",
"laion2B-s32B-b79K",
precision=self.dtype,
device=self.device,
jit=False,
force_quick_gelu=False,
force_custom_text=False,
force_patch_dropout=False,
force_image_size=None,
pretrained_image=False,
image_mean=None,
image_std=None,
light_augmentation=True,
aug_cfg={},
output_dict=True,
with_score_predictor=False,
with_region_predictor=False,
)
self.tokenizer = get_tokenizer("ViT-H-14")
# https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/blob/main/preprocessor_config.json
# TODO: [transforms.Resize(224), transforms.CenterCrop(224)] for any aspect ratio.
self.transform = transforms.Compose([
transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BICUBIC),
transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]),
])
if version == "v2.0":
url = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Third_Party/HPS_v2_compressed.pt"
filename = "HPS_v2_compressed.pt"
md5 = "fd9180de357abf01fdb4eaad64631db4"
elif version == "v2.1":
url = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Third_Party/HPS_v2.1_compressed.pt"
filename = "HPS_v2.1_compressed.pt"
md5 = "4067542e34ba2553a738c5ac6c1d75c0"
else:
raise ValueError("Only v2.0 and v2.1 are supported.")
if self.model_path is None or not os.path.exists(self.model_path):
download_url(url, torch.hub.get_dir(), md5=md5)
model_path = os.path.join(torch.hub.get_dir(), filename)
state_dict = torch.load(model_path, map_location="cpu")["state_dict"]
self.model.load_state_dict(state_dict)
self.model.to(device=self.device, dtype=self.dtype)
self.model.requires_grad_(False)
self.model.eval()
def __call__(self, batch_frames: torch.Tensor, batch_prompt: list[str]) -> Tuple[torch.Tensor, torch.Tensor]:
assert batch_frames.shape[0] == len(batch_prompt)
# Compute batch reward and loss in frame-wise.
batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w")
batch_loss, batch_reward = 0, 0
for frames in batch_frames:
image_inputs = torch.stack([self.transform(frame) for frame in frames])
image_inputs = image_inputs.to(device=self.device, dtype=self.dtype)
text_inputs = self.tokenizer(batch_prompt).to(device=self.device)
outputs = self.model(image_inputs, text_inputs)
image_features, text_features = outputs["image_features"], outputs["text_features"]
logits = image_features @ text_features.T
reward = torch.diagonal(logits)
# Convert reward to loss in [0, 1].
if self.max_reward is None:
loss = (-1 * reward) * self.loss_scale
else:
loss = abs(reward - self.max_reward) * self.loss_scale
batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean()
return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0]
class PickScoreReward(BaseReward):
"""[PickScore](https://github.com/yuvalkirstain/PickScore) reward model.
"""
def __init__(
self,
model_path="yuvalkirstain/PickScore_v1",
device="cpu",
dtype=torch.float16,
max_reward=1,
loss_scale=1,
):
from transformers import AutoProcessor, AutoModel
self.model_path = model_path
self.device = device
self.dtype = dtype
self.max_reward = max_reward
self.loss_scale = loss_scale
# https://huggingface.co/yuvalkirstain/PickScore_v1/blob/main/preprocessor_config.json
self.transform = transforms.Compose([
transforms.Resize(224, interpolation=transforms.InterpolationMode.BICUBIC),
transforms.CenterCrop(224),
transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]),
])
self.processor = AutoProcessor.from_pretrained("laion/CLIP-ViT-H-14-laion2B-s32B-b79K", torch_dtype=self.dtype)
self.model = AutoModel.from_pretrained(model_path, torch_dtype=self.dtype).eval().to(device)
self.model.requires_grad_(False)
self.model.eval()
def __call__(self, batch_frames: torch.Tensor, batch_prompt: list[str]) -> Tuple[torch.Tensor, torch.Tensor]:
assert batch_frames.shape[0] == len(batch_prompt)
# Compute batch reward and loss in frame-wise.
batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w")
batch_loss, batch_reward = 0, 0
for frames in batch_frames:
image_inputs = torch.stack([self.transform(frame) for frame in frames])
image_inputs = image_inputs.to(device=self.device, dtype=self.dtype)
text_inputs = self.processor(
text=batch_prompt,
padding=True,
truncation=True,
max_length=77,
return_tensors="pt",
).to(self.device)
image_features = self.model.get_image_features(pixel_values=image_inputs)
text_features = self.model.get_text_features(**text_inputs)
image_features = image_features / torch.norm(image_features, dim=-1, keepdim=True)
text_features = text_features / torch.norm(text_features, dim=-1, keepdim=True)
logits = image_features @ text_features.T
reward = torch.diagonal(logits)
# Convert reward to loss in [0, 1].
if self.max_reward is None:
loss = (-1 * reward) * self.loss_scale
else:
loss = abs(reward - self.max_reward) * self.loss_scale
batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean()
return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0]
class MPSReward(BaseReward):
"""[MPS](https://github.com/Kwai-Kolors/MPS) reward model.
"""
def __init__(
self,
model_path=None,
device="cpu",
dtype=torch.float16,
max_reward=1,
loss_scale=1,
):
from transformers import AutoTokenizer, AutoConfig
from .MPS.trainer.models.clip_model import CLIPModel
self.model_path = model_path
self.device = device
self.dtype = dtype
self.condition = "light, color, clarity, tone, style, ambiance, artistry, shape, face, hair, hands, limbs, structure, instance, texture, quantity, attributes, position, number, location, word, things."
self.max_reward = max_reward
self.loss_scale = loss_scale
processor_name_or_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K"
# https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/blob/main/preprocessor_config.json
# TODO: [transforms.Resize(224), transforms.CenterCrop(224)] for any aspect ratio.
self.transform = transforms.Compose([
transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BICUBIC),
transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]),
])
# We convert the original [ckpt](http://drive.google.com/file/d/17qrK_aJkVNM75ZEvMEePpLj6L867MLkN/view?usp=sharing)
# (contains the entire model) to a `state_dict`.
url = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Third_Party/MPS_overall.pth"
filename = "MPS_overall.pth"
md5 = "1491cbbbd20565747fe07e7572e2ac56"
if self.model_path is None or not os.path.exists(self.model_path):
download_url(url, torch.hub.get_dir(), md5=md5)
model_path = os.path.join(torch.hub.get_dir(), filename)
self.tokenizer = AutoTokenizer.from_pretrained(processor_name_or_path, trust_remote_code=True)
config = AutoConfig.from_pretrained(processor_name_or_path)
self.model = CLIPModel(config)
state_dict = torch.load(model_path, map_location="cpu")
self.model.load_state_dict(state_dict, strict=False)
self.model.to(device=self.device, dtype=self.dtype)
self.model.requires_grad_(False)
self.model.eval()
def _tokenize(self, caption):
input_ids = self.tokenizer(
caption,
max_length=self.tokenizer.model_max_length,
padding="max_length",
truncation=True,
return_tensors="pt"
).input_ids
return input_ids
def __call__(
self,
batch_frames: torch.Tensor,
batch_prompt: list[str],
batch_condition: Optional[list[str]] = None
) -> Tuple[torch.Tensor, torch.Tensor]:
if batch_condition is None:
batch_condition = [self.condition] * len(batch_prompt)
batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w")
batch_loss, batch_reward = 0, 0
for frames in batch_frames:
image_inputs = torch.stack([self.transform(frame) for frame in frames])
image_inputs = image_inputs.to(device=self.device, dtype=self.dtype)
text_inputs = self._tokenize(batch_prompt).to(self.device)
condition_inputs = self._tokenize(batch_condition).to(device=self.device)
text_features, image_features = self.model(text_inputs, image_inputs, condition_inputs)
text_features = text_features / text_features.norm(dim=-1, keepdim=True)
image_features = image_features / image_features.norm(dim=-1, keepdim=True)
# reward = self.model.logit_scale.exp() * torch.diag(torch.einsum('bd,cd->bc', text_features, image_features))
logits = image_features @ text_features.T
reward = torch.diagonal(logits)
# Convert reward to loss in [0, 1].
if self.max_reward is None:
loss = (-1 * reward) * self.loss_scale
else:
loss = abs(reward - self.max_reward) * self.loss_scale
batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean()
return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0]
if __name__ == "__main__":
import numpy as np
from decord import VideoReader
video_path_list = ["your_video_path_1.mp4", "your_video_path_2.mp4"]
prompt_list = ["your_prompt_1", "your_prompt_2"]
num_sampled_frames = 8
to_tensor = transforms.ToTensor()
sampled_frames_list = []
for video_path in video_path_list:
vr = VideoReader(video_path)
sampled_frame_indices = np.linspace(0, len(vr), num_sampled_frames, endpoint=False, dtype=int)
sampled_frames = vr.get_batch(sampled_frame_indices).asnumpy()
sampled_frames = torch.stack([to_tensor(frame) for frame in sampled_frames])
sampled_frames_list.append(sampled_frames)
sampled_frames = torch.stack(sampled_frames_list)
sampled_frames = rearrange(sampled_frames, "b t c h w -> b c t h w")
aesthetic_reward_v2 = AestheticReward(device="cuda", dtype=torch.bfloat16)
print(f"aesthetic_reward_v2: {aesthetic_reward_v2(sampled_frames)}")
aesthetic_reward_v2_5 = AestheticReward(
encoder_path="google/siglip-so400m-patch14-384", version="v2.5", device="cuda", dtype=torch.bfloat16
)
print(f"aesthetic_reward_v2_5: {aesthetic_reward_v2_5(sampled_frames)}")
hps_reward_v2 = HPSReward(device="cuda", dtype=torch.bfloat16)
print(f"hps_reward_v2: {hps_reward_v2(sampled_frames, prompt_list)}")
hps_reward_v2_1 = HPSReward(version="v2.1", device="cuda", dtype=torch.bfloat16)
print(f"hps_reward_v2_1: {hps_reward_v2_1(sampled_frames, prompt_list)}")
pick_score = PickScoreReward(device="cuda", dtype=torch.bfloat16)
print(f"pick_score_reward: {pick_score(sampled_frames, prompt_list)}")
mps_score = MPSReward(device="cuda", dtype=torch.bfloat16)
print(f"mps_reward: {mps_score(sampled_frames, prompt_list)}")
+1812 -109
View File
File diff suppressed because it is too large Load Diff
-51
View File
@@ -1,51 +0,0 @@
# Modified from OpenAI's diffusion repos
# GLIDE: https://github.com/openai/glide-text2im/blob/main/glide_text2im/gaussian_diffusion.py
# ADM: https://github.com/openai/guided-diffusion/blob/main/guided_diffusion
# IDDPM: https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py
from . import gaussian_diffusion as gd
from .respace import SpacedDiffusion, space_timesteps
def IDDPM(
timestep_respacing,
noise_schedule="linear",
use_kl=False,
sigma_small=False,
predict_xstart=False,
learn_sigma=True,
pred_sigma=True,
rescale_learned_sigmas=False,
diffusion_steps=1000,
snr=False,
return_startx=False,
):
betas = gd.get_named_beta_schedule(noise_schedule, diffusion_steps)
if use_kl:
loss_type = gd.LossType.RESCALED_KL
elif rescale_learned_sigmas:
loss_type = gd.LossType.RESCALED_MSE
else:
loss_type = gd.LossType.MSE
if timestep_respacing is None or timestep_respacing == "":
timestep_respacing = [diffusion_steps]
return SpacedDiffusion(
use_timesteps=space_timesteps(diffusion_steps, timestep_respacing),
betas=betas,
model_mean_type=(
gd.ModelMeanType.START_X if predict_xstart else gd.ModelMeanType.EPSILON
),
model_var_type=(
(gd.ModelVarType.LEARNED_RANGE if learn_sigma else (
gd.ModelVarType.FIXED_LARGE
if not sigma_small
else gd.ModelVarType.FIXED_SMALL
)
)
if pred_sigma
else None
),
loss_type=loss_type,
snr=snr,
return_startx=return_startx,
# rescale_timesteps=rescale_timesteps,
)
+46
View File
@@ -0,0 +1,46 @@
"""Modified from https://github.com/THUDM/CogVideo/blob/3710a612d8760f5cdb1741befeebb65b9e0f2fe0/sat/sgm/modules/diffusionmodules/sigma_sampling.py
"""
import torch
class DiscreteSampling:
def __init__(self, num_idx, uniform_sampling=False):
self.num_idx = num_idx
self.uniform_sampling = uniform_sampling
self.is_distributed = torch.distributed.is_available() and torch.distributed.is_initialized()
if self.is_distributed and self.uniform_sampling:
world_size = torch.distributed.get_world_size()
self.rank = torch.distributed.get_rank()
i = 1
while True:
if world_size % i != 0 or num_idx % (world_size // i) != 0:
i += 1
else:
self.group_num = world_size // i
break
assert self.group_num > 0
assert world_size % self.group_num == 0
# the number of rank in one group
self.group_width = world_size // self.group_num
self.sigma_interval = self.num_idx // self.group_num
print('rank=%d world_size=%d group_num=%d group_width=%d sigma_interval=%s' % (
self.rank, world_size, self.group_num,
self.group_width, self.sigma_interval))
def __call__(self, n_samples, generator=None, device=None):
if self.is_distributed and self.uniform_sampling:
group_index = self.rank // self.group_width
idx = torch.randint(
group_index * self.sigma_interval,
(group_index + 1) * self.sigma_interval,
(n_samples,),
generator=generator, device=device,
)
print('proc[%d] idx=%s' % (self.rank, idx))
else:
idx = torch.randint(
0, self.num_idx, (n_samples,),
generator=generator, device=device,
)
return idx
+28
View File
@@ -0,0 +1,28 @@
"""Modified from https://github.com/kijai/ComfyUI-MochiWrapper
"""
import torch
import torch.nn as nn
def autocast_model_forward(cls, origin_dtype, *inputs, **kwargs):
weight_dtype = cls.weight.dtype
cls.to(origin_dtype)
# Convert all inputs to the original dtype
inputs = [input.to(origin_dtype) for input in inputs]
out = cls.original_forward(*inputs, **kwargs)
cls.to(weight_dtype)
return out
def convert_weight_dtype_wrapper(module, origin_dtype):
for name, module in module.named_modules():
if name == "":
continue
original_forward = module.forward
if hasattr(module, "weight"):
setattr(module, "original_forward", original_forward)
setattr(
module,
"forward",
lambda *inputs, m=module, **kwargs: autocast_model_forward(m, origin_dtype, *inputs, **kwargs)
)
+65 -47
View File
@@ -156,8 +156,8 @@ def precalculate_safetensors_hashes(tensors, metadata):
class LoRANetwork(torch.nn.Module):
TRANSFORMER_TARGET_REPLACE_MODULE = ["Transformer2DModel", "Transformer3DModel"]
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["T5LayerSelfAttention", "T5LayerFF"]
TRANSFORMER_TARGET_REPLACE_MODULE = ["Transformer2DModel", "Transformer3DModel", "HunyuanTransformer3DModel", "EasyAnimateTransformer3DModel"]
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["T5LayerSelfAttention", "T5LayerFF", "BertEncoder"]
LORA_PREFIX_TRANSFORMER = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
def __init__(
@@ -238,9 +238,10 @@ class LoRANetwork(torch.nn.Module):
self.text_encoder_loras = []
skipped_te = []
for i, text_encoder in enumerate(text_encoders):
text_encoder_loras, skipped = create_modules(False, text_encoder, LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE)
self.text_encoder_loras.extend(text_encoder_loras)
skipped_te += skipped
if text_encoder is not None:
text_encoder_loras, skipped = create_modules(False, text_encoder, LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE)
self.text_encoder_loras.extend(text_encoder_loras)
skipped_te += skipped
print(f"create LoRA for Text Encoder: {len(self.text_encoder_loras)} modules.")
self.unet_loras, skipped_un = create_modules(True, unet, LoRANetwork.TRANSFORMER_TARGET_REPLACE_MODULE)
@@ -389,36 +390,44 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
layer_infos = layer.split(LORA_PREFIX_TRANSFORMER + "_")[-1].split("_")
curr_layer = pipeline.transformer
temp_name = layer_infos.pop(0)
while len(layer_infos) > -1:
try:
curr_layer = curr_layer.__getattr__(temp_name)
if len(layer_infos) > 0:
temp_name = layer_infos.pop(0)
elif len(layer_infos) == 0:
break
except Exception:
if len(layer_infos) == 0:
print('Error loading layer')
if len(temp_name) > 0:
temp_name += "_" + layer_infos.pop(0)
else:
temp_name = layer_infos.pop(0)
try:
curr_layer = curr_layer.__getattr__("_".join(layer_infos[1:]))
except Exception:
temp_name = layer_infos.pop(0)
while len(layer_infos) > -1:
try:
curr_layer = curr_layer.__getattr__(temp_name)
if len(layer_infos) > 0:
temp_name = layer_infos.pop(0)
elif len(layer_infos) == 0:
break
except Exception:
if len(layer_infos) == 0:
print('Error loading layer')
if len(temp_name) > 0:
temp_name += "_" + layer_infos.pop(0)
else:
temp_name = layer_infos.pop(0)
weight_up = elems['lora_up.weight'].to(dtype)
weight_down = elems['lora_down.weight'].to(dtype)
origin_dtype = curr_layer.weight.data.dtype
origin_device = curr_layer.weight.data.device
curr_layer = curr_layer.to(device, dtype)
weight_up = elems['lora_up.weight'].to(device, dtype)
weight_down = elems['lora_down.weight'].to(device, dtype)
if 'alpha' in elems.keys():
alpha = elems['alpha'].item() / weight_up.shape[1]
else:
alpha = 1.0
curr_layer.weight.data = curr_layer.weight.data.to(device)
if len(weight_up.shape) == 4:
curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up.squeeze(3).squeeze(2),
weight_down.squeeze(3).squeeze(2)).unsqueeze(
2).unsqueeze(3)
curr_layer.weight.data += multiplier * alpha * torch.mm(
weight_up.squeeze(3).squeeze(2), weight_down.squeeze(3).squeeze(2)
).unsqueeze(2).unsqueeze(3)
else:
curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up, weight_down)
curr_layer = curr_layer.to(origin_device, origin_dtype)
return pipeline
@@ -443,34 +452,43 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl
layer_infos = layer.split(LORA_PREFIX_UNET + "_")[-1].split("_")
curr_layer = pipeline.transformer
temp_name = layer_infos.pop(0)
while len(layer_infos) > -1:
try:
curr_layer = curr_layer.__getattr__(temp_name)
if len(layer_infos) > 0:
temp_name = layer_infos.pop(0)
elif len(layer_infos) == 0:
break
except Exception:
if len(layer_infos) == 0:
print('Error loading layer')
if len(temp_name) > 0:
temp_name += "_" + layer_infos.pop(0)
else:
temp_name = layer_infos.pop(0)
try:
curr_layer = curr_layer.__getattr__("_".join(layer_infos[1:]))
except Exception:
temp_name = layer_infos.pop(0)
while len(layer_infos) > -1:
try:
curr_layer = curr_layer.__getattr__(temp_name)
if len(layer_infos) > 0:
temp_name = layer_infos.pop(0)
elif len(layer_infos) == 0:
break
except Exception:
if len(layer_infos) == 0:
print('Error loading layer')
if len(temp_name) > 0:
temp_name += "_" + layer_infos.pop(0)
else:
temp_name = layer_infos.pop(0)
weight_up = elems['lora_up.weight'].to(dtype)
weight_down = elems['lora_down.weight'].to(dtype)
origin_dtype = curr_layer.weight.data.dtype
origin_device = curr_layer.weight.data.device
curr_layer = curr_layer.to(device, dtype)
weight_up = elems['lora_up.weight'].to(device, dtype)
weight_down = elems['lora_down.weight'].to(device, dtype)
if 'alpha' in elems.keys():
alpha = elems['alpha'].item() / weight_up.shape[1]
else:
alpha = 1.0
curr_layer.weight.data = curr_layer.weight.data.to(device)
if len(weight_up.shape) == 4:
curr_layer.weight.data -= multiplier * alpha * torch.mm(weight_up.squeeze(3).squeeze(2),
weight_down.squeeze(3).squeeze(2)).unsqueeze(2).unsqueeze(3)
curr_layer.weight.data -= multiplier * alpha * torch.mm(
weight_up.squeeze(3).squeeze(2), weight_down.squeeze(3).squeeze(2)
).unsqueeze(2).unsqueeze(3)
else:
curr_layer.weight.data -= multiplier * alpha * torch.mm(weight_up, weight_down)
curr_layer = curr_layer.to(origin_device, origin_dtype)
return pipeline
return pipeline
+185 -1
View File
@@ -1,5 +1,7 @@
import gc
import os
import cv2
import imageio
import numpy as np
import torch
@@ -8,7 +10,43 @@ from einops import rearrange
from PIL import Image
def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=6, fps=12, imageio_backend=True):
def get_width_and_height_from_image_and_base_resolution(image, base_resolution):
target_pixels = int(base_resolution) * int(base_resolution)
original_width, original_height = Image.open(image).size
ratio = (target_pixels / (original_width * original_height)) ** 0.5
width_slider = round(original_width * ratio)
height_slider = round(original_height * ratio)
return height_slider, width_slider
def color_transfer(sc, dc):
"""
Transfer color distribution from of sc, referred to dc.
Args:
sc (numpy.ndarray): input image to be transfered.
dc (numpy.ndarray): reference image
Returns:
numpy.ndarray: Transferred color distribution on the sc.
"""
def get_mean_and_std(img):
x_mean, x_std = cv2.meanStdDev(img)
x_mean = np.hstack(np.around(x_mean, 2))
x_std = np.hstack(np.around(x_std, 2))
return x_mean, x_std
sc = cv2.cvtColor(sc, cv2.COLOR_RGB2LAB)
s_mean, s_std = get_mean_and_std(sc)
dc = cv2.cvtColor(dc, cv2.COLOR_RGB2LAB)
t_mean, t_std = get_mean_and_std(dc)
img_n = ((sc - s_mean) * (t_std / s_std)) + t_mean
np.putmask(img_n, img_n > 255, 255)
np.putmask(img_n, img_n < 0, 0)
dst = cv2.cvtColor(cv2.convertScaleAbs(img_n), cv2.COLOR_LAB2RGB)
return dst
def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=6, fps=12, imageio_backend=True, color_transfer_post_process=False):
videos = rearrange(videos, "b c t h w -> t b c h w")
outputs = []
for x in videos:
@@ -19,6 +57,10 @@ def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=6, f
x = (x * 255).numpy().astype(np.uint8)
outputs.append(Image.fromarray(x))
if color_transfer_post_process:
for i in range(1, len(outputs)):
outputs[i] = Image.fromarray(color_transfer(np.uint8(outputs[i]), np.uint8(outputs[0])))
os.makedirs(os.path.dirname(path), exist_ok=True)
if imageio_backend:
if path.endswith("mp4"):
@@ -29,3 +71,145 @@ def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=6, f
if path.endswith("mp4"):
path = path.replace('.mp4', '.gif')
outputs[0].save(path, format='GIF', append_images=outputs, save_all=True, duration=100, loop=0)
def get_image_to_video_latent(validation_image_start, validation_image_end, video_length, sample_size):
if validation_image_start is not None and validation_image_end is not None:
if type(validation_image_start) is str and os.path.isfile(validation_image_start):
image_start = clip_image = Image.open(validation_image_start).convert("RGB")
image_start = image_start.resize([sample_size[1], sample_size[0]])
clip_image = clip_image.resize([sample_size[1], sample_size[0]])
else:
image_start = clip_image = validation_image_start
image_start = [_image_start.resize([sample_size[1], sample_size[0]]) for _image_start in image_start]
clip_image = [_clip_image.resize([sample_size[1], sample_size[0]]) for _clip_image in clip_image]
if type(validation_image_end) is str and os.path.isfile(validation_image_end):
image_end = Image.open(validation_image_end).convert("RGB")
image_end = image_end.resize([sample_size[1], sample_size[0]])
else:
image_end = validation_image_end
image_end = [_image_end.resize([sample_size[1], sample_size[0]]) for _image_end in image_end]
if type(image_start) is list:
clip_image = clip_image[0]
start_video = torch.cat(
[torch.from_numpy(np.array(_image_start)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0) for _image_start in image_start],
dim=2
)
input_video = torch.tile(start_video[:, :, :1], [1, 1, video_length, 1, 1])
input_video[:, :, :len(image_start)] = start_video
input_video_mask = torch.zeros_like(input_video[:, :1])
input_video_mask[:, :, len(image_start):] = 255
else:
input_video = torch.tile(
torch.from_numpy(np.array(image_start)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0),
[1, 1, video_length, 1, 1]
)
input_video_mask = torch.zeros_like(input_video[:, :1])
input_video_mask[:, :, 1:] = 255
if type(image_end) is list:
image_end = [_image_end.resize(image_start[0].size if type(image_start) is list else image_start.size) for _image_end in image_end]
end_video = torch.cat(
[torch.from_numpy(np.array(_image_end)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0) for _image_end in image_end],
dim=2
)
input_video[:, :, -len(end_video):] = end_video
input_video_mask[:, :, -len(image_end):] = 0
else:
image_end = image_end.resize(image_start[0].size if type(image_start) is list else image_start.size)
input_video[:, :, -1:] = torch.from_numpy(np.array(image_end)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0)
input_video_mask[:, :, -1:] = 0
input_video = input_video / 255
elif validation_image_start is not None:
if type(validation_image_start) is str and os.path.isfile(validation_image_start):
image_start = clip_image = Image.open(validation_image_start).convert("RGB")
image_start = image_start.resize([sample_size[1], sample_size[0]])
clip_image = clip_image.resize([sample_size[1], sample_size[0]])
else:
image_start = clip_image = validation_image_start
image_start = [_image_start.resize([sample_size[1], sample_size[0]]) for _image_start in image_start]
clip_image = [_clip_image.resize([sample_size[1], sample_size[0]]) for _clip_image in clip_image]
image_end = None
if type(image_start) is list:
clip_image = clip_image[0]
start_video = torch.cat(
[torch.from_numpy(np.array(_image_start)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0) for _image_start in image_start],
dim=2
)
input_video = torch.tile(start_video[:, :, :1], [1, 1, video_length, 1, 1])
input_video[:, :, :len(image_start)] = start_video
input_video = input_video / 255
input_video_mask = torch.zeros_like(input_video[:, :1])
input_video_mask[:, :, len(image_start):] = 255
else:
input_video = torch.tile(
torch.from_numpy(np.array(image_start)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0),
[1, 1, video_length, 1, 1]
) / 255
input_video_mask = torch.zeros_like(input_video[:, :1])
input_video_mask[:, :, 1:, ] = 255
else:
image_start = None
image_end = None
input_video = torch.zeros([1, 3, video_length, sample_size[0], sample_size[1]])
input_video_mask = torch.ones([1, 1, video_length, sample_size[0], sample_size[1]]) * 255
clip_image = None
del image_start
del image_end
gc.collect()
return input_video, input_video_mask, clip_image
def get_video_to_video_latent(input_video_path, video_length, sample_size, fps=None, validation_video_mask=None, ref_image=None):
if isinstance(input_video_path, str):
cap = cv2.VideoCapture(input_video_path)
input_video = []
original_fps = cap.get(cv2.CAP_PROP_FPS)
frame_skip = 1 if fps is None else int(original_fps // fps)
frame_count = 0
while True:
ret, frame = cap.read()
if not ret:
break
if frame_count % frame_skip == 0:
frame = cv2.resize(frame, (sample_size[1], sample_size[0]))
input_video.append(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB))
frame_count += 1
cap.release()
else:
input_video = input_video_path
input_video = torch.from_numpy(np.array(input_video))[:video_length]
input_video = input_video.permute([3, 0, 1, 2]).unsqueeze(0) / 255
if ref_image is not None:
ref_image = Image.open(ref_image)
ref_image = torch.from_numpy(np.array(ref_image))
ref_image = ref_image.unsqueeze(0).permute([3, 0, 1, 2]).unsqueeze(0) / 255
if validation_video_mask is not None:
validation_video_mask = Image.open(validation_video_mask).convert('L').resize((sample_size[1], sample_size[0]))
input_video_mask = np.where(np.array(validation_video_mask) < 240, 0, 255)
input_video_mask = torch.from_numpy(np.array(input_video_mask)).unsqueeze(0).unsqueeze(-1).permute([3, 0, 1, 2]).unsqueeze(0)
input_video_mask = torch.tile(input_video_mask, [1, 1, input_video.size()[2], 1, 1])
input_video_mask = input_video_mask.to(input_video.device, input_video.dtype)
else:
input_video_mask = torch.zeros_like(input_video[:, :1])
input_video_mask[:, :, :] = 255
return input_video, input_video_mask, ref_image
+82
View File
@@ -0,0 +1,82 @@
Copyright (c) 2022 Robin Rombach and Patrick Esser and contributors
CreativeML Open RAIL-M
dated August 22, 2022
Section I: PREAMBLE
Multimodal generative models are being widely adopted and used, and have the potential to transform the way artists, among other individuals, conceive and benefit from AI or ML technologies as a tool for content creation.
Notwithstanding the current and potential benefits that these artifacts can bring to society at large, there are also concerns about potential misuses of them, either due to their technical limitations or ethical considerations.
In short, this license strives for both the open and responsible downstream use of the accompanying model. When it comes to the open character, we took inspiration from open source permissive licenses regarding the grant of IP rights. Referring to the downstream responsible use, we added use-based restrictions not permitting the use of the Model in very specific scenarios, in order for the licensor to be able to enforce the license in case potential misuses of the Model may occur. At the same time, we strive to promote open and responsible research on generative models for art and content generation.
Even though downstream derivative versions of the model could be released under different licensing terms, the latter will always have to include - at minimum - the same use-based restrictions as the ones in the original license (this license). We believe in the intersection between open and responsible AI development; thus, this License aims to strike a balance between both in order to enable responsible open-science in the field of AI.
This License governs the use of the model (and its derivatives) and is informed by the model card associated with the model.
NOW THEREFORE, You and Licensor agree as follows:
1. Definitions
- "License" means the terms and conditions for use, reproduction, and Distribution as defined in this document.
- "Data" means a collection of information and/or content extracted from the dataset used with the Model, including to train, pretrain, or otherwise evaluate the Model. The Data is not licensed under this License.
- "Output" means the results of operating a Model as embodied in informational content resulting therefrom.
- "Model" means any accompanying machine-learning based assemblies (including checkpoints), consisting of learnt weights, parameters (including optimizer states), corresponding to the model architecture as embodied in the Complementary Material, that have been trained or tuned, in whole or in part on the Data, using the Complementary Material.
- "Derivatives of the Model" means all modifications to the Model, works based on the Model, or any other model which is created or initialized by transfer of patterns of the weights, parameters, activations or output of the Model, to the other model, in order to cause the other model to perform similarly to the Model, including - but not limited to - distillation methods entailing the use of intermediate data representations or methods based on the generation of synthetic data by the Model for training the other model.
- "Complementary Material" means the accompanying source code and scripts used to define, run, load, benchmark or evaluate the Model, and used to prepare data for training or evaluation, if any. This includes any accompanying documentation, tutorials, examples, etc, if any.
- "Distribution" means any transmission, reproduction, publication or other sharing of the Model or Derivatives of the Model to a third party, including providing the Model as a hosted service made available by electronic or other remote means - e.g. API-based or web access.
- "Licensor" means the copyright owner or entity authorized by the copyright owner that is granting the License, including the persons or entities that may have rights in the Model and/or distributing the Model.
- "You" (or "Your") means an individual or Legal Entity exercising permissions granted by this License and/or making use of the Model for whichever purpose and in any field of use, including usage of the Model in an end-use application - e.g. chatbot, translator, image generator.
- "Third Parties" means individuals or legal entities that are not under common control with Licensor or You.
- "Contribution" means any work of authorship, including the original version of the Model and any modifications or additions to that Model or Derivatives of the Model thereof, that is intentionally submitted to Licensor for inclusion in the Model by the copyright owner or by an individual or Legal Entity authorized to submit on behalf of the copyright owner. For the purposes of this definition, "submitted" means any form of electronic, verbal, or written communication sent to the Licensor or its representatives, including but not limited to communication on electronic mailing lists, source code control systems, and issue tracking systems that are managed by, or on behalf of, the Licensor for the purpose of discussing and improving the Model, but excluding communication that is conspicuously marked or otherwise designated in writing by the copyright owner as "Not a Contribution."
- "Contributor" means Licensor and any individual or Legal Entity on behalf of whom a Contribution has been received by Licensor and subsequently incorporated within the Model.
Section II: INTELLECTUAL PROPERTY RIGHTS
Both copyright and patent grants apply to the Model, Derivatives of the Model and Complementary Material. The Model and Derivatives of the Model are subject to additional terms as described in Section III.
2. Grant of Copyright License. Subject to the terms and conditions of this License, each Contributor hereby grants to You a perpetual, worldwide, non-exclusive, no-charge, royalty-free, irrevocable copyright license to reproduce, prepare, publicly display, publicly perform, sublicense, and distribute the Complementary Material, the Model, and Derivatives of the Model.
3. Grant of Patent License. Subject to the terms and conditions of this License and where and as applicable, each Contributor hereby grants to You a perpetual, worldwide, non-exclusive, no-charge, royalty-free, irrevocable (except as stated in this paragraph) patent license to make, have made, use, offer to sell, sell, import, and otherwise transfer the Model and the Complementary Material, where such license applies only to those patent claims licensable by such Contributor that are necessarily infringed by their Contribution(s) alone or by combination of their Contribution(s) with the Model to which such Contribution(s) was submitted. If You institute patent litigation against any entity (including a cross-claim or counterclaim in a lawsuit) alleging that the Model and/or Complementary Material or a Contribution incorporated within the Model and/or Complementary Material constitutes direct or contributory patent infringement, then any patent licenses granted to You under this License for the Model and/or Work shall terminate as of the date such litigation is asserted or filed.
Section III: CONDITIONS OF USAGE, DISTRIBUTION AND REDISTRIBUTION
4. Distribution and Redistribution. You may host for Third Party remote access purposes (e.g. software-as-a-service), reproduce and distribute copies of the Model or Derivatives of the Model thereof in any medium, with or without modifications, provided that You meet the following conditions:
Use-based restrictions as referenced in paragraph 5 MUST be included as an enforceable provision by You in any type of legal agreement (e.g. a license) governing the use and/or distribution of the Model or Derivatives of the Model, and You shall give notice to subsequent users You Distribute to, that the Model or Derivatives of the Model are subject to paragraph 5. This provision does not apply to the use of Complementary Material.
You must give any Third Party recipients of the Model or Derivatives of the Model a copy of this License;
You must cause any modified files to carry prominent notices stating that You changed the files;
You must retain all copyright, patent, trademark, and attribution notices excluding those notices that do not pertain to any part of the Model, Derivatives of the Model.
You may add Your own copyright statement to Your modifications and may provide additional or different license terms and conditions - respecting paragraph 4.a. - for use, reproduction, or Distribution of Your modifications, or for any such Derivatives of the Model as a whole, provided Your use, reproduction, and Distribution of the Model otherwise complies with the conditions stated in this License.
5. Use-based restrictions. The restrictions set forth in Attachment A are considered Use-based restrictions. Therefore You cannot use the Model and the Derivatives of the Model for the specified restricted uses. You may use the Model subject to this License, including only for lawful purposes and in accordance with the License. Use may include creating any content with, finetuning, updating, running, training, evaluating and/or reparametrizing the Model. You shall require all of Your users who use the Model or a Derivative of the Model to comply with the terms of this paragraph (paragraph 5).
6. The Output You Generate. Except as set forth herein, Licensor claims no rights in the Output You generate using the Model. You are accountable for the Output you generate and its subsequent uses. No use of the output can contravene any provision as stated in the License.
Section IV: OTHER PROVISIONS
7. Updates and Runtime Restrictions. To the maximum extent permitted by law, Licensor reserves the right to restrict (remotely or otherwise) usage of the Model in violation of this License, update the Model through electronic means, or modify the Output of the Model based on updates. You shall undertake reasonable efforts to use the latest version of the Model.
8. Trademarks and related. Nothing in this License permits You to make use of Licensors’ trademarks, trade names, logos or to otherwise suggest endorsement or misrepresent the relationship between the parties; and any rights not expressly granted herein are reserved by the Licensors.
9. Disclaimer of Warranty. Unless required by applicable law or agreed to in writing, Licensor provides the Model and the Complementary Material (and each Contributor provides its Contributions) on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied, including, without limitation, any warranties or conditions of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A PARTICULAR PURPOSE. You are solely responsible for determining the appropriateness of using or redistributing the Model, Derivatives of the Model, and the Complementary Material and assume any risks associated with Your exercise of permissions under this License.
10. Limitation of Liability. In no event and under no legal theory, whether in tort (including negligence), contract, or otherwise, unless required by applicable law (such as deliberate and grossly negligent acts) or agreed to in writing, shall any Contributor be liable to You for damages, including any direct, indirect, special, incidental, or consequential damages of any character arising as a result of this License or out of the use or inability to use the Model and the Complementary Material (including but not limited to damages for loss of goodwill, work stoppage, computer failure or malfunction, or any and all other commercial damages or losses), even if such Contributor has been advised of the possibility of such damages.
11. Accepting Warranty or Additional Liability. While redistributing the Model, Derivatives of the Model and the Complementary Material thereof, You may choose to offer, and charge a fee for, acceptance of support, warranty, indemnity, or other liability obligations and/or rights consistent with this License. However, in accepting such obligations, You may act only on Your own behalf and on Your sole responsibility, not on behalf of any other Contributor, and only if You agree to indemnify, defend, and hold each Contributor harmless for any liability incurred by, or claims asserted against, such Contributor by reason of your accepting any such warranty or additional liability.
12. If any provision of this License is held to be invalid, illegal or unenforceable, the remaining provisions shall be unaffected thereby and remain valid as if such provision had not been set forth herein.
END OF TERMS AND CONDITIONS
Attachment A
Use Restrictions
You agree not to use the Model or Derivatives of the Model:
- In any way that violates any applicable national, federal, state, local or international law or regulation;
- For the purpose of exploiting, harming or attempting to exploit or harm minors in any way;
- To generate or disseminate verifiably false information and/or content with the purpose of harming others;
- To generate or disseminate personal identifiable information that can be used to harm an individual;
- To defame, disparage or otherwise harass others;
- For fully automated decision making that adversely impacts an individual’s legal rights or otherwise creates or modifies a binding, enforceable obligation;
- For any use intended to or which has the effect of discriminating against or harming individuals or groups based on online or offline social behavior or known or predicted personal or personality characteristics;
- To exploit any of the vulnerabilities of a specific group of persons based on their age, social, physical or mental characteristics, in order to materially distort the behavior of a person pertaining to that group in a manner that causes or is likely to cause that person or another person physical or psychological harm;
- For any use intended to or which has the effect of discriminating against individuals or groups based on legally protected characteristics or categories;
- To provide medical advice and medical results interpretation;
- To generate or disseminate information for the purpose to be used for administration of justice, law enforcement, immigration or asylum processes, such as predicting an individual will commit fraud/crime commitment (e.g. by text profiling, drawing causal relationships between assertions made in documents, indiscriminate and arbitrarily-targeted use).
+63
View File
@@ -0,0 +1,63 @@
## VAE Training
English | [简体中文](./README_zh-CN.md)
After completing data preprocessing, we can obtain the following dataset:
```
📦 project/
├── 📂 datasets/
│ ├── 📂 internal_datasets/
│ ├── 📂 videos/
│ │ ├── 📄 00000001.mp4
│ │ ├── 📄 00000001.jpg
│ │ └── 📄 .....
│ └── 📄 json_of_internal_datasets.json
```
The json_of_internal_datasets.json is a standard JSON file. The file_path in the json can to be set as relative path, as shown in below:
```json
[
{
"file_path": "videos/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
{
"file_path": "train/00000001.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
.....
]
```
You can also set the path as absolute path as follow:
```json
[
{
"file_path": "/mnt/data/videos/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
{
"file_path": "/mnt/data/train/00000001.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
.....
]
```
## Train Video VAE
We need to set config in ```easyanimate/vae/configs/autoencoder``` at first. The default config is ```autoencoder_kl_32x32x4_slice.yaml```. We need to set the some params in yaml file.
- ```data_json_path``` corresponds to the JSON file of the dataset.
- ```data_root``` corresponds to the root path of the dataset. If you want to use absolute path in json file, please delete this line.
- ```ckpt_path``` corresponds to the pretrained weights of the vae.
- ```gpus``` and num_nodes need to be set as the actual situation of your machine.
The we run shell file as follow:
```
sh scripts/train_vae.sh
```
+63
View File
@@ -0,0 +1,63 @@
## VAE 训练
[English](./README.md) | 简体中文
在完成数据预处理后,你可以获得这样的数据格式:
```
📦 project/
├── 📂 datasets/
│ ├── 📂 internal_datasets/
│ ├── 📂 videos/
│ │ ├── 📄 00000001.mp4
│ │ ├── 📄 00000001.jpg
│ │ └── 📄 .....
│ └── 📄 json_of_internal_datasets.json
```
json_of_internal_datasets.json是一个标准的json文件。json中的file_path可以被设置为相对路径,如下所示:
```json
[
{
"file_path": "videos/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
{
"file_path": "train/00000001.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
.....
]
```
你也可以将路径设置为绝对路径:
```json
[
{
"file_path": "/mnt/data/videos/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
{
"file_path": "/mnt/data/train/00000001.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
.....
]
```
## 训练 Video VAE
我们首先需要修改 ```easyanimate/vae/configs/autoencoder``` 中的配置文件。默认的配置文件是 ```autoencoder_kl_32x32x4_slice.yaml```。你需要修改以下参数:
- ```data_json_path``` json file 所在的目录。
- ```data_root``` 数据的根目录。如果你在json file中使用了绝对路径,请设置为空。
- ```ckpt_path``` 预训练的vae模型路径。
- ```gpus``` 以及 ```num_nodes``` 需要设置为你机器的实际gpu数目。
运行以下的脚本来训练vae:
```
sh scripts/train_vae.sh
```
@@ -0,0 +1,64 @@
model:
base_learning_rate: 1.0e-04
target: easyanimate.vae.ldm.models.cogvideox_casual3dcnn.AutoencoderKLMagvit_CogVideoX
params:
latent_channels: 16
temporal_compression_ratio: 4
monitor: train/rec_loss
ckpt_path: vae/diffusion_pytorch_model.safetensors
down_block_types: ("CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D",
"CogVideoXDownBlock3D",)
up_block_types: ("CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D",
"CogVideoXUpBlock3D",)
lossconfig:
target: easyanimate.vae.ldm.modules.losses.LPIPSWithDiscriminator
params:
disc_start: 50001
kl_weight: 1.0e-06
disc_weight: 0.5
l2_loss_weight: 0.1
l1_loss_weight: 1.0
perceptual_weight: 1.0
data:
target: train_vae.DataModuleFromConfig
params:
batch_size: 1
wrap: true
num_workers: 8
train:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRTrain
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 256
degradation: pil_nearest
video_size: 256
video_len: 49
slice_interval: 1
validation:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRValidation
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 256
degradation: pil_nearest
video_size: 256
video_len: 49
slice_interval: 1
lightning:
callbacks:
image_logger:
target: train_vae.ImageLogger
params:
batch_frequency: 5000
max_images: 8
increase_log_steps: True
trainer:
benchmark: True
accumulate_grad_batches: 1
gpus: "0"
num_nodes: 1
@@ -0,0 +1,62 @@
model:
base_learning_rate: 1.0e-04
target: easyanimate.vae.ldm.models.omnigen_casual3dcnn.AutoencoderKLMagvit_fromOmnigen
params:
monitor: train/rec_loss
ckpt_path: models/videoVAE_omnigen_8x8x4_from_vae-ft-mse-840000-ema-pruned.ckpt
down_block_types: ("SpatialDownBlock3D", "SpatialTemporalDownBlock3D", "SpatialTemporalDownBlock3D",
"SpatialTemporalDownBlock3D",)
up_block_types: ("SpatialUpBlock3D", "SpatialTemporalUpBlock3D", "SpatialTemporalUpBlock3D",
"SpatialTemporalUpBlock3D",)
lossconfig:
target: easyanimate.vae.ldm.modules.losses.LPIPSWithDiscriminator
params:
disc_start: 50001
kl_weight: 1.0e-06
disc_weight: 0.5
l2_loss_weight: 0.1
l1_loss_weight: 1.0
perceptual_weight: 1.0
data:
target: train_vae.DataModuleFromConfig
params:
batch_size: 2
wrap: true
num_workers: 4
train:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRTrain
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 128
degradation: pil_nearest
video_size: 128
video_len: 9
slice_interval: 1
validation:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRValidation
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 128
degradation: pil_nearest
video_size: 128
video_len: 9
slice_interval: 1
lightning:
callbacks:
image_logger:
target: train_vae.ImageLogger
params:
batch_frequency: 5000
max_images: 8
increase_log_steps: True
trainer:
benchmark: True
accumulate_grad_batches: 1
gpus: "0"
num_nodes: 1
@@ -0,0 +1,65 @@
model:
base_learning_rate: 1.0e-04
target: easyanimate.vae.ldm.models.omnigen_casual3dcnn.AutoencoderKLMagvit_fromOmnigen
params:
spatial_group_norm: true
mid_block_attention_type: "spatial"
latent_channels: 16
monitor: train/rec_loss
ckpt_path: vae/diffusion_pytorch_model.safetensors
down_block_types: ("SpatialDownBlock3D", "SpatialTemporalDownBlock3D", "SpatialTemporalDownBlock3D",
"SpatialTemporalDownBlock3D",)
up_block_types: ("SpatialUpBlock3D", "SpatialTemporalUpBlock3D", "SpatialTemporalUpBlock3D",
"SpatialTemporalUpBlock3D",)
lossconfig:
target: easyanimate.vae.ldm.modules.losses.LPIPSWithDiscriminator
params:
disc_start: 50001
kl_weight: 1.0e-06
disc_weight: 0.5
l2_loss_weight: 0.1
l1_loss_weight: 1.0
perceptual_weight: 1.0
data:
target: train_vae.DataModuleFromConfig
params:
batch_size: 1
wrap: true
num_workers: 8
train:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRTrain
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 256
degradation: pil_nearest
video_size: 256
video_len: 49
slice_interval: 1
validation:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRValidation
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 256
degradation: pil_nearest
video_size: 256
video_len: 49
slice_interval: 1
lightning:
callbacks:
image_logger:
target: train_vae.ImageLogger
params:
batch_frequency: 5000
max_images: 8
increase_log_steps: True
trainer:
benchmark: True
accumulate_grad_batches: 1
gpus: "0"
num_nodes: 1
@@ -0,0 +1,65 @@
model:
base_learning_rate: 1.0e-04
target: easyanimate.vae.ldm.models.omnigen_casual3dcnn.AutoencoderKLMagvit_fromOmnigen
params:
slice_compression_vae: true
mini_batch_encoder: 8
mini_batch_decoder: 2
monitor: train/rec_loss
ckpt_path: models/Diffusion_Transformer/EasyAnimateV2-XL-2-512x512/vae/diffusion_pytorch_model.safetensors
down_block_types: ("SpatialDownBlock3D", "SpatialTemporalDownBlock3D", "SpatialTemporalDownBlock3D",
"SpatialTemporalDownBlock3D",)
up_block_types: ("SpatialUpBlock3D", "SpatialTemporalUpBlock3D", "SpatialTemporalUpBlock3D",
"SpatialTemporalUpBlock3D",)
lossconfig:
target: easyanimate.vae.ldm.modules.losses.LPIPSWithDiscriminator
params:
disc_start: 50001
kl_weight: 1.0e-06
disc_weight: 0.5
l2_loss_weight: 0.0
l1_loss_weight: 1.0
perceptual_weight: 1.0
data:
target: train_vae.DataModuleFromConfig
params:
batch_size: 1
wrap: true
num_workers: 8
train:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRTrain
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 256
degradation: pil_nearest
video_size: 256
video_len: 25
slice_interval: 1
validation:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRValidation
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 256
degradation: pil_nearest
video_size: 256
video_len: 25
slice_interval: 1
lightning:
callbacks:
image_logger:
target: train_vae.ImageLogger
params:
batch_frequency: 5000
max_images: 8
increase_log_steps: True
trainer:
benchmark: True
accumulate_grad_batches: 1
gpus: "0"
num_nodes: 1
@@ -0,0 +1,66 @@
model:
base_learning_rate: 1.0e-04
target: easyanimate.vae.ldm.models.omnigen_casual3dcnn.AutoencoderKLMagvit_fromOmnigen
params:
slice_compression_vae: true
train_decoder_only: true
mini_batch_encoder: 8
mini_batch_decoder: 2
monitor: train/rec_loss
ckpt_path: models/Diffusion_Transformer/EasyAnimateV2-XL-2-512x512/vae/diffusion_pytorch_model.safetensors
down_block_types: ("SpatialDownBlock3D", "SpatialTemporalDownBlock3D", "SpatialTemporalDownBlock3D",
"SpatialTemporalDownBlock3D",)
up_block_types: ("SpatialUpBlock3D", "SpatialTemporalUpBlock3D", "SpatialTemporalUpBlock3D",
"SpatialTemporalUpBlock3D",)
lossconfig:
target: easyanimate.vae.ldm.modules.losses.LPIPSWithDiscriminator
params:
disc_start: 50001
kl_weight: 1.0e-06
disc_weight: 0.5
l2_loss_weight: 1.0
l1_loss_weight: 0.0
perceptual_weight: 1.0
data:
target: train_vae.DataModuleFromConfig
params:
batch_size: 1
wrap: true
num_workers: 8
train:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRTrain
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 256
degradation: pil_nearest
video_size: 256
video_len: 25
slice_interval: 1
validation:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRValidation
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 256
degradation: pil_nearest
video_size: 256
video_len: 25
slice_interval: 1
lightning:
callbacks:
image_logger:
target: train_vae.ImageLogger
params:
batch_frequency: 5000
max_images: 8
increase_log_steps: True
trainer:
benchmark: True
accumulate_grad_batches: 1
gpus: "0"
num_nodes: 1
@@ -0,0 +1,66 @@
model:
base_learning_rate: 1.0e-04
target: easyanimate.vae.ldm.models.omnigen_casual3dcnn.AutoencoderKLMagvit_fromOmnigen
params:
slice_compression_vae: true
mini_batch_encoder: 8
mini_batch_decoder: 1
monitor: train/rec_loss
ckpt_path: models/Diffusion_Transformer/EasyAnimateV2-XL-2-512x512/vae/diffusion_pytorch_model.safetensors
down_block_types: ("SpatialTemporalDownBlock3D", "SpatialTemporalDownBlock3D", "SpatialTemporalDownBlock3D",
"SpatialTemporalDownBlock3D",)
up_block_types: ("SpatialTemporalUpBlock3D", "SpatialTemporalUpBlock3D", "SpatialTemporalUpBlock3D",
"SpatialTemporalUpBlock3D",)
lossconfig:
target: easyanimate.vae.ldm.modules.losses.LPIPSWithDiscriminator
params:
disc_start: 50001
kl_weight: 1.0e-06
disc_weight: 0.5
l2_loss_weight: 0.0
l1_loss_weight: 1.0
perceptual_weight: 1.0
data:
target: train_vae.DataModuleFromConfig
params:
batch_size: 1
wrap: true
num_workers: 8
train:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRTrain
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 256
degradation: pil_nearest
video_size: 256
video_len: 33
slice_interval: 1
validation:
target: easyanimate.vae.ldm.data.dataset_image_video.CustomSRValidation
params:
data_json_path: pretrain.json
data_root: /your_data_root # This is used in relative path
size: 256
degradation: pil_nearest
video_size: 256
video_len: 33
slice_interval: 1
lightning:
callbacks:
image_logger:
target: train_vae.ImageLogger
params:
batch_frequency: 5000
max_images: 8
increase_log_steps: True
trainer:
benchmark: True
accumulate_grad_batches: 1
gpus: "0"
num_nodes: 1
+29
View File
@@ -0,0 +1,29 @@
name: ldm
channels:
- pytorch
- defaults
dependencies:
- python=3.8.5
- pip=20.3
- cudatoolkit=11.3
- pytorch=1.11.0
- torchvision=0.12.0
- numpy=1.19.2
- pip:
- albumentations==0.4.3
- diffusers
- opencv-python==4.1.2.30
- pudb==2019.2
- invisible-watermark
- imageio==2.9.0
- imageio-ffmpeg==0.4.2
- pytorch-lightning==1.4.2
- omegaconf==2.1.1
- test-tube>=0.7.5
- streamlit>=0.73.1
- einops==0.3.0
- torch-fidelity==0.3.0
- transformers==4.19.2
- torchmetrics==0.6.0
- kornia==0.6
- -e .
+25
View File
@@ -0,0 +1,25 @@
from abc import abstractmethod
from torch.utils.data import (ChainDataset, ConcatDataset, Dataset,
IterableDataset)
class Txt2ImgIterableBaseDataset(IterableDataset):
'''
Define an interface to make the IterableDatasets for text2img data chainable
'''
def __init__(self, num_records=0, valid_ids=None, size=256):
super().__init__()
self.num_records = num_records
self.valid_ids = valid_ids
self.sample_ids = valid_ids
self.size = size
print(f'{self.__class__.__name__} dataset contains {self.__len__()} examples.')
def __len__(self):
return self.num_records
@abstractmethod
def __iter__(self):
pass
@@ -0,0 +1,26 @@
#-*- encoding:utf-8 -*-
from pytorch_lightning.callbacks import Callback
class DatasetCallback(Callback):
def __init__(self):
self.sampler_pos_start = 0
self.preload_used_idx_flag = False
def on_train_start(self, trainer, pl_module):
if not self.preload_used_idx_flag:
self.preload_used_idx_flag = True
trainer.train_dataloader.batch_sampler.sampler_pos_reload = self.sampler_pos_start
def on_save_checkpoint(self, trainer, pl_module, checkpoint):
if trainer.train_dataloader is not None:
# Save sampler_pos_start parameters in the checkpoint
checkpoint['sampler_pos_start'] = trainer.train_dataloader.batch_sampler.sampler_pos_start
def on_load_checkpoint(self, trainer, pl_module, checkpoint):
# Restore sampler_pos_start parameters from the checkpoint
if 'sampler_pos_start' in checkpoint:
self.sampler_pos_start = checkpoint.get('sampler_pos_start', 0)
print('Load sampler_pos_start from checkpoint, sampler_pos_start = %d' % self.sampler_pos_start)
else:
print('The sampler_pos_start is not in checkpoint')
@@ -0,0 +1,284 @@
import glob
import json
import os
import pickle
import random
import shutil
import tarfile
from functools import partial
import albumentations
import cv2
import numpy as np
import PIL
import torchvision.transforms.functional as TF
import yaml
from decord import VideoReader
from func_timeout import FunctionTimedOut, func_set_timeout
from omegaconf import OmegaConf
from PIL import Image
from torch.utils.data import BatchSampler, Dataset, Sampler
from tqdm import tqdm
from ..modules.image_degradation import (degradation_fn_bsr,
degradation_fn_bsr_light)
class ImageVideoSampler(BatchSampler):
"""A sampler wrapper for grouping images with similar aspect ratio into a same batch.
Args:
sampler (Sampler): Base sampler.
dataset (Dataset): Dataset providing data information.
batch_size (int): Size of mini-batch.
drop_last (bool): If ``True``, the sampler will drop the last batch if
its size would be less than ``batch_size``.
aspect_ratios (dict): The predefined aspect ratios.
"""
def __init__(self,
sampler: Sampler,
dataset: Dataset,
batch_size: int,
drop_last: bool = False
) -> None:
if not isinstance(sampler, Sampler):
raise TypeError('sampler should be an instance of ``Sampler``, '
f'but got {sampler}')
if not isinstance(batch_size, int) or batch_size <= 0:
raise ValueError('batch_size should be a positive integer value, '
f'but got batch_size={batch_size}')
self.sampler = sampler
self.dataset = dataset
self.batch_size = batch_size
self.drop_last = drop_last
self.sampler_pos_start = 0
self.sampler_pos_reload = 0
self.num_samples_random = len(self.sampler)
# buckets for each aspect ratio
self.bucket = {'image':[], 'video':[]}
def set_epoch(self, epoch):
if hasattr(self.sampler, "set_epoch"):
self.sampler.set_epoch(epoch)
def __iter__(self):
for index_sampler, idx in enumerate(self.sampler):
if self.sampler_pos_reload != 0 and self.sampler_pos_reload < self.num_samples_random:
if index_sampler < self.sampler_pos_reload:
self.sampler_pos_start = (self.sampler_pos_start + 1) % self.num_samples_random
continue
elif index_sampler == self.sampler_pos_reload:
self.sampler_pos_reload = 0
content_type = self.dataset.data.get_type(idx)
bucket = self.bucket[content_type]
bucket.append(idx)
# yield a batch of indices in the same aspect ratio group
if len(self.bucket['video']) == self.batch_size:
yield self.bucket['video']
self.bucket['video'] = []
elif len(self.bucket['image']) == self.batch_size:
yield self.bucket['image']
self.bucket['image'] = []
self.sampler_pos_start = (self.sampler_pos_start + 1) % self.num_samples_random
class ImageVideoDataset(Dataset):
# update __getitem__() from ImageNetSR. If timeout for Pandas70M, throw exception.
# If caught exception(timeout or others), try another index until successful and return.
def __init__(self, size=None, video_size=128, video_len=25,
degradation=None, downscale_f=4, random_crop=True, min_crop_f=0.25, max_crop_f=1.,
s_t=None, slice_interval=None, data_root=None
):
"""
Imagenet Superresolution Dataloader
Performs following ops in order:
1. crops a crop of size s from image either as random or center crop
2. resizes crop to size with cv2.area_interpolation
3. degrades resized crop with degradation_fn
:param size: resizing to size after cropping
:param degradation: degradation_fn, e.g. cv_bicubic or bsrgan_light
:param downscale_f: Low Resolution Downsample factor
:param min_crop_f: determines crop size s,
where s = c * min_img_side_len with c sampled from interval (min_crop_f, max_crop_f)
:param max_crop_f: ""
:param data_root:
:param random_crop:
"""
self.base = self.get_base()
assert size
assert (size / downscale_f).is_integer()
self.size = size
self.LR_size = int(size / downscale_f)
self.min_crop_f = min_crop_f
self.max_crop_f = max_crop_f
assert(max_crop_f <= 1.)
self.center_crop = not random_crop
self.s_t = s_t
self.slice_interval = slice_interval
self.image_rescaler = albumentations.SmallestMaxSize(max_size=size, interpolation=cv2.INTER_AREA)
self.video_rescaler = albumentations.SmallestMaxSize(max_size=video_size, interpolation=cv2.INTER_AREA)
self.video_len = video_len
self.video_size = video_size
self.data_root = data_root
self.pil_interpolation = False # gets reset later if incase interp_op is from pillow
if degradation == "bsrgan":
self.degradation_process = partial(degradation_fn_bsr, sf=downscale_f)
elif degradation == "bsrgan_light":
self.degradation_process = partial(degradation_fn_bsr_light, sf=downscale_f)
else:
interpolation_fn = {
"cv_nearest": cv2.INTER_NEAREST,
"cv_bilinear": cv2.INTER_LINEAR,
"cv_bicubic": cv2.INTER_CUBIC,
"cv_area": cv2.INTER_AREA,
"cv_lanczos": cv2.INTER_LANCZOS4,
"pil_nearest": PIL.Image.NEAREST,
"pil_bilinear": PIL.Image.BILINEAR,
"pil_bicubic": PIL.Image.BICUBIC,
"pil_box": PIL.Image.BOX,
"pil_hamming": PIL.Image.HAMMING,
"pil_lanczos": PIL.Image.LANCZOS,
}[degradation]
self.pil_interpolation = degradation.startswith("pil_")
if self.pil_interpolation:
self.degradation_process = partial(TF.resize, size=self.LR_size, interpolation=interpolation_fn)
else:
self.degradation_process = albumentations.SmallestMaxSize(max_size=self.LR_size,
interpolation=interpolation_fn)
def __len__(self):
return len(self.base)
def get_type(self, index):
return self.base[index].get('type', 'image')
def __getitem__(self, i):
@func_set_timeout(15) # time wait 3 seconds
def get_video_item(example):
if self.data_root is not None:
video_reader = VideoReader(os.path.join(self.data_root, example['file_path']))
else:
video_reader = VideoReader(example['file_path'])
video_length = len(video_reader)
if self.slice_interval == "rand":
slice_interval = np.random.choice([1, 2, 3])
else:
slice_interval = int(self.slice_interval)
clip_length = min(video_length, (self.video_len - 1) * slice_interval + 1)
start_idx = random.randint(0, video_length - clip_length)
batch_index = np.linspace(start_idx, start_idx + clip_length - 1, self.video_len, dtype=int)
pixel_values = video_reader.get_batch(batch_index).asnumpy()
del video_reader
out_images = []
LR_out_images = []
min_side_len = min(pixel_values[0].shape[:2])
crop_side_len = min_side_len * np.random.uniform(self.min_crop_f, self.max_crop_f, size=None)
crop_side_len = int(crop_side_len)
if self.center_crop:
self.cropper = albumentations.CenterCrop(height=crop_side_len, width=crop_side_len)
else:
self.cropper = albumentations.RandomCrop(height=crop_side_len, width=crop_side_len)
imgs = np.transpose(pixel_values, (1, 2, 3, 0))
imgs = self.cropper(image=imgs)["image"]
imgs = np.transpose(imgs, (3, 0, 1, 2))
for img in imgs:
image = self.video_rescaler(image=img)["image"]
out_images.append(image[None, :, :, :])
if self.pil_interpolation:
image_pil = PIL.Image.fromarray(image)
LR_image = self.degradation_process(image_pil)
LR_image = np.array(LR_image).astype(np.uint8)
else:
LR_image = self.degradation_process(image=image)["image"]
LR_out_images.append(LR_image[None, :, :, :])
example = {}
example['image'] = (np.concatenate(out_images) / 127.5 - 1.0).astype(np.float32)
example['LR_image'] = (np.concatenate(LR_out_images) / 127.5 - 1.0).astype(np.float32)
return example
example = self.base[i]
if example.get('type', 'image') == 'video':
while True:
try:
example = self.base[i]
return get_video_item(example)
except FunctionTimedOut:
print("stt catch: Function 'extract failed' timed out.")
i = random.randint(0, self.__len__() - 1)
except Exception as e:
print('stt catch', e)
i = random.randint(0, self.__len__() - 1)
elif example.get('type', 'image') == 'image':
while True:
try:
example = self.base[i]
if self.data_root is not None:
image = Image.open(os.path.join(self.data_root, example['file_path']))
else:
image = Image.open(example['file_path'])
image = image.convert("RGB")
image = np.array(image).astype(np.uint8)
min_side_len = min(image.shape[:2])
crop_side_len = min_side_len * np.random.uniform(self.min_crop_f, self.max_crop_f, size=None)
crop_side_len = int(crop_side_len)
if self.center_crop:
self.cropper = albumentations.CenterCrop(height=crop_side_len, width=crop_side_len)
else:
self.cropper = albumentations.RandomCrop(height=crop_side_len, width=crop_side_len)
image = self.cropper(image=image)["image"]
image = self.image_rescaler(image=image)["image"]
if self.pil_interpolation:
image_pil = PIL.Image.fromarray(image)
LR_image = self.degradation_process(image_pil)
LR_image = np.array(LR_image).astype(np.uint8)
else:
LR_image = self.degradation_process(image=image)["image"]
example = {}
example["image"] = (image/127.5 - 1.0).astype(np.float32)
example["LR_image"] = (LR_image/127.5 - 1.0).astype(np.float32)
return example
except Exception as e:
print("catch", e)
i = random.randint(0, self.__len__() - 1)
class CustomSRTrain(ImageVideoDataset):
def __init__(self, data_json_path, **kwargs):
self.data_json_path = data_json_path
super().__init__(**kwargs)
def get_base(self):
return [ann for ann in json.load(open(self.data_json_path))]
class CustomSRValidation(ImageVideoDataset):
def __init__(self, data_json_path, **kwargs):
self.data_json_path = data_json_path
super().__init__(**kwargs)
self.data_json_path = data_json_path
def get_base(self):
return [ann for ann in json.load(open(self.data_json_path))][:100] + \
[ann for ann in json.load(open(self.data_json_path))][-100:]
+98
View File
@@ -0,0 +1,98 @@
import numpy as np
class LambdaWarmUpCosineScheduler:
"""
note: use with a base_lr of 1.0
"""
def __init__(self, warm_up_steps, lr_min, lr_max, lr_start, max_decay_steps, verbosity_interval=0):
self.lr_warm_up_steps = warm_up_steps
self.lr_start = lr_start
self.lr_min = lr_min
self.lr_max = lr_max
self.lr_max_decay_steps = max_decay_steps
self.last_lr = 0.
self.verbosity_interval = verbosity_interval
def schedule(self, n, **kwargs):
if self.verbosity_interval > 0:
if n % self.verbosity_interval == 0: print(f"current step: {n}, recent lr-multiplier: {self.last_lr}")
if n < self.lr_warm_up_steps:
lr = (self.lr_max - self.lr_start) / self.lr_warm_up_steps * n + self.lr_start
self.last_lr = lr
return lr
else:
t = (n - self.lr_warm_up_steps) / (self.lr_max_decay_steps - self.lr_warm_up_steps)
t = min(t, 1.0)
lr = self.lr_min + 0.5 * (self.lr_max - self.lr_min) * (
1 + np.cos(t * np.pi))
self.last_lr = lr
return lr
def __call__(self, n, **kwargs):
return self.schedule(n,**kwargs)
class LambdaWarmUpCosineScheduler2:
"""
supports repeated iterations, configurable via lists
note: use with a base_lr of 1.0.
"""
def __init__(self, warm_up_steps, f_min, f_max, f_start, cycle_lengths, verbosity_interval=0):
assert len(warm_up_steps) == len(f_min) == len(f_max) == len(f_start) == len(cycle_lengths)
self.lr_warm_up_steps = warm_up_steps
self.f_start = f_start
self.f_min = f_min
self.f_max = f_max
self.cycle_lengths = cycle_lengths
self.cum_cycles = np.cumsum([0] + list(self.cycle_lengths))
self.last_f = 0.
self.verbosity_interval = verbosity_interval
def find_in_interval(self, n):
interval = 0
for cl in self.cum_cycles[1:]:
if n <= cl:
return interval
interval += 1
def schedule(self, n, **kwargs):
cycle = self.find_in_interval(n)
n = n - self.cum_cycles[cycle]
if self.verbosity_interval > 0:
if n % self.verbosity_interval == 0: print(f"current step: {n}, recent lr-multiplier: {self.last_f}, "
f"current cycle {cycle}")
if n < self.lr_warm_up_steps[cycle]:
f = (self.f_max[cycle] - self.f_start[cycle]) / self.lr_warm_up_steps[cycle] * n + self.f_start[cycle]
self.last_f = f
return f
else:
t = (n - self.lr_warm_up_steps[cycle]) / (self.cycle_lengths[cycle] - self.lr_warm_up_steps[cycle])
t = min(t, 1.0)
f = self.f_min[cycle] + 0.5 * (self.f_max[cycle] - self.f_min[cycle]) * (
1 + np.cos(t * np.pi))
self.last_f = f
return f
def __call__(self, n, **kwargs):
return self.schedule(n, **kwargs)
class LambdaLinearScheduler(LambdaWarmUpCosineScheduler2):
def schedule(self, n, **kwargs):
cycle = self.find_in_interval(n)
n = n - self.cum_cycles[cycle]
if self.verbosity_interval > 0:
if n % self.verbosity_interval == 0: print(f"current step: {n}, recent lr-multiplier: {self.last_f}, "
f"current cycle {cycle}")
if n < self.lr_warm_up_steps[cycle]:
f = (self.f_max[cycle] - self.f_start[cycle]) / self.lr_warm_up_steps[cycle] * n + self.f_start[cycle]
self.last_f = f
return f
else:
f = self.f_min[cycle] + (self.f_max[cycle] - self.f_min[cycle]) * (self.cycle_lengths[cycle] - n) / (self.cycle_lengths[cycle])
self.last_f = f
return f
+337
View File
@@ -0,0 +1,337 @@
import time
from contextlib import contextmanager
import pytorch_lightning as pl
import torch
import torch.nn.functional as F
from ..modules.diffusionmodules.model import Decoder, Encoder
from ..modules.distributions.distributions import DiagonalGaussianDistribution
from ..util import instantiate_from_config
from .enc_dec_pytorch import Decoder as Mag_Decoder
from .enc_dec_pytorch import Encoder as Mag_Encoder
class AutoencoderKLMagvit(pl.LightningModule):
def __init__(self,
ddconfig,
lossconfig,
embed_dim,
ckpt_path=None,
ignore_keys=[],
image_key="image",
colorize_nlabels=None,
monitor=None,
):
super().__init__()
self.image_key = image_key
self.encoder = Mag_Encoder()
self.decoder = Mag_Decoder()
self.loss = instantiate_from_config(lossconfig)
self.quant_conv = torch.nn.Conv3d(16, 16, 1)
self.post_quant_conv = torch.nn.Conv3d(8, 8, 1)
self.embed_dim = embed_dim
if colorize_nlabels is not None:
assert type(colorize_nlabels)==int
self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1))
if monitor is not None:
self.monitor = monitor
if ckpt_path is not None:
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)
def init_from_ckpt(self, path, ignore_keys=list()):
sd = torch.load(path, map_location="cpu")["state_dict"]
keys = list(sd.keys())
for k in keys:
for ik in ignore_keys:
if k.startswith(ik):
print("Deleting key {} from state_dict.".format(k))
del sd[k]
self.load_state_dict(sd, strict=False)
print(f"Restored from {path}")
def encode(self, x):
h = self.encoder(x)
moments = self.quant_conv(h)
posterior = DiagonalGaussianDistribution(moments)
return posterior
def decode(self, z):
z = self.post_quant_conv(z)
dec = self.decoder(z)
return dec
def forward(self, input, sample_posterior=True):
if input.ndim==4:
input = input.unsqueeze(2)
posterior = self.encode(input)
if sample_posterior:
z = posterior.sample()
else:
z = posterior.mode()
dec = self.decode(z)
return dec, posterior
def get_input(self, batch, k):
x = batch[k]
if x.ndim==5:
x = x.permute(0, 4, 1, 2, 3).to(memory_format=torch.contiguous_format).float()
return x
if len(x.shape) == 3:
x = x[..., None]
x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format).float()
return x
def training_step(self, batch, batch_idx, optimizer_idx):
# tic = time.time()
inputs = self.get_input(batch, self.image_key)
# print(f"get_input time {time.time() - tic}")
# tic = time.time()
reconstructions, posterior = self(inputs)
# print(f"model forward time {time.time() - tic}")
if optimizer_idx == 0:
# train encoder+decoder+logvar
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("aeloss", aeloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=False)
# print(f"cal loss time {time.time() - tic}")
return aeloss
if optimizer_idx == 1:
# train the discriminator
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("discloss", discloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=False)
# print(f"cal loss time {time.time() - tic}")
return discloss
def validation_step(self, batch, batch_idx):
with torch.no_grad():
inputs = self.get_input(batch, self.image_key)
reconstructions, posterior = self(inputs)
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, 0, self.global_step,
last_layer=self.get_last_layer(), split="val")
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, 1, self.global_step,
last_layer=self.get_last_layer(), split="val")
self.log("val/rec_loss", log_dict_ae["val/rec_loss"])
self.log_dict(log_dict_ae)
self.log_dict(log_dict_disc)
return self.log_dict
def configure_optimizers(self):
lr = self.learning_rate
opt_ae = torch.optim.Adam(list(self.encoder.parameters())+
list(self.decoder.parameters())+
list(self.quant_conv.parameters())+
list(self.post_quant_conv.parameters()),
lr=lr, betas=(0.5, 0.9))
opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(),
lr=lr, betas=(0.5, 0.9))
return [opt_ae, opt_disc], []
def get_last_layer(self):
return self.decoder.conv_out.weight
@torch.no_grad()
def log_images(self, batch, only_inputs=False, **kwargs):
log = dict()
x = self.get_input(batch, self.image_key)
x = x.to(self.device)
if not only_inputs:
xrec, posterior = self(x)
if x.shape[1] > 3:
# colorize with random projection
assert xrec.shape[1] > 3
x = self.to_rgb(x)
xrec = self.to_rgb(xrec)
log["samples"] = self.decode(torch.randn_like(posterior.sample()))
log["reconstructions"] = xrec
log["inputs"] = x
return log
def to_rgb(self, x):
assert self.image_key == "segmentation"
if not hasattr(self, "colorize"):
self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x))
x = F.conv2d(x, weight=self.colorize)
x = 2.*(x-x.min())/(x.max()-x.min()) - 1.
return x
class AutoencoderKL(pl.LightningModule):
def __init__(self,
ddconfig,
lossconfig,
embed_dim,
ckpt_path=None,
ignore_keys=[],
image_key="image",
colorize_nlabels=None,
monitor=None,
):
super().__init__()
self.image_key = image_key
self.encoder = Encoder(**ddconfig)
self.decoder = Decoder(**ddconfig)
self.loss = instantiate_from_config(lossconfig)
assert ddconfig["double_z"]
self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1)
self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
self.embed_dim = embed_dim
if colorize_nlabels is not None:
assert type(colorize_nlabels)==int
self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1))
if monitor is not None:
self.monitor = monitor
if ckpt_path is not None:
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)
def init_from_ckpt(self, path, ignore_keys=list()):
sd = torch.load(path, map_location="cpu")["state_dict"]
keys = list(sd.keys())
for k in keys:
for ik in ignore_keys:
if k.startswith(ik):
print("Deleting key {} from state_dict.".format(k))
del sd[k]
self.load_state_dict(sd, strict=False)
print(f"Restored from {path}")
def encode(self, x):
h = self.encoder(x)
moments = self.quant_conv(h)
posterior = DiagonalGaussianDistribution(moments)
return posterior
def decode(self, z):
z = self.post_quant_conv(z)
dec = self.decoder(z)
return dec
def forward(self, input, sample_posterior=True):
posterior = self.encode(input)
if sample_posterior:
z = posterior.sample()
else:
z = posterior.mode()
dec = self.decode(z)
return dec, posterior
def get_input(self, batch, k):
x = batch[k]
if len(x.shape) == 3:
x = x[..., None]
x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format).float()
return x
def training_step(self, batch, batch_idx, optimizer_idx):
# tic = time.time()
inputs = self.get_input(batch, self.image_key)
# print(f"get_input time {time.time() - tic}")
# tic = time.time()
reconstructions, posterior = self(inputs)
# print(f"model forward time {time.time() - tic}")
tic = time.time()
if optimizer_idx == 0:
# train encoder+decoder+logvar
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("aeloss", aeloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=False)
# print(f"cal loss time {time.time() - tic}")
return aeloss
if optimizer_idx == 1:
# train the discriminator
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("discloss", discloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=False)
# print(f"cal loss time {time.time() - tic}")
return discloss
def validation_step(self, batch, batch_idx):
tic = time.time()
inputs = self.get_input(batch, self.image_key)
print(f"get_input time {time.time() - tic}")
tic = time.time()
reconstructions, posterior = self(inputs)
print(f"val forward time {time.time() - tic}")
tic = time.time()
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, 0, self.global_step,
last_layer=self.get_last_layer(), split="val")
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, 1, self.global_step,
last_layer=self.get_last_layer(), split="val")
self.log("val/rec_loss", log_dict_ae["val/rec_loss"])
self.log_dict(log_dict_ae)
self.log_dict(log_dict_disc)
print(f"val end time {time.time() - tic}")
return self.log_dict
def configure_optimizers(self):
lr = self.learning_rate
opt_ae = torch.optim.Adam(list(self.encoder.parameters())+
list(self.decoder.parameters())+
list(self.quant_conv.parameters())+
list(self.post_quant_conv.parameters()),
lr=lr, betas=(0.5, 0.9))
opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(),
lr=lr, betas=(0.5, 0.9))
return [opt_ae, opt_disc], []
def get_last_layer(self):
return self.decoder.conv_out.weight
@torch.no_grad()
def log_images(self, batch, only_inputs=False, **kwargs):
log = dict()
x = self.get_input(batch, self.image_key)
x = x.to(self.device)
if not only_inputs:
xrec, posterior = self(x)
if x.shape[1] > 3:
# colorize with random projection
assert xrec.shape[1] > 3
x = self.to_rgb(x)
xrec = self.to_rgb(xrec)
log["samples"] = self.decode(torch.randn_like(posterior.sample()))
log["reconstructions"] = xrec
log["inputs"] = x
return log
def to_rgb(self, x):
assert self.image_key == "segmentation"
if not hasattr(self, "colorize"):
self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x))
x = F.conv2d(x, weight=self.colorize)
x = 2.*(x-x.min())/(x.max()-x.min()) - 1.
return x
class IdentityFirstStage(torch.nn.Module):
def __init__(self, *args, vq_interface=False, **kwargs):
self.vq_interface = vq_interface # TODO: Should be true by default but check to not break older stuff
super().__init__()
def encode(self, x, *args, **kwargs):
return x
def decode(self, x, *args, **kwargs):
return x
def quantize(self, x, *args, **kwargs):
if self.vq_interface:
return x, None, [None, None, None]
return x
def forward(self, x, *args, **kwargs):
return x
+337
View File
@@ -0,0 +1,337 @@
import time
from contextlib import contextmanager
import pytorch_lightning as pl
import torch
import torch.nn.functional as F
from ..modules.diffusionmodules.model import Decoder, Encoder
from ..modules.distributions.distributions import DiagonalGaussianDistribution
from ..util import instantiate_from_config
from .enc_dec import Decoder as Mag_Decoder
from .enc_dec import Encoder as Mag_Encoder
class AutoencoderKLMagvit(pl.LightningModule):
def __init__(self,
ddconfig,
lossconfig,
embed_dim,
ckpt_path=None,
ignore_keys=[],
image_key="image",
colorize_nlabels=None,
monitor=None,
):
super().__init__()
self.image_key = image_key
self.encoder = Mag_Encoder()
self.decoder = Mag_Decoder()
self.loss = instantiate_from_config(lossconfig)
self.quant_conv = torch.nn.Conv3d(16, 16, 1)
self.post_quant_conv = torch.nn.Conv3d(8, 8, 1)
self.embed_dim = embed_dim
if colorize_nlabels is not None:
assert type(colorize_nlabels)==int
self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1))
if monitor is not None:
self.monitor = monitor
if ckpt_path is not None:
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)
def init_from_ckpt(self, path, ignore_keys=list()):
sd = torch.load(path, map_location="cpu")["state_dict"]
keys = list(sd.keys())
for k in keys:
for ik in ignore_keys:
if k.startswith(ik):
print("Deleting key {} from state_dict.".format(k))
del sd[k]
self.load_state_dict(sd, strict=False)
print(f"Restored from {path}")
def encode(self, x):
h = self.encoder(x)
moments = self.quant_conv(h)
posterior = DiagonalGaussianDistribution(moments)
return posterior
def decode(self, z):
z = self.post_quant_conv(z)
dec = self.decoder(z)
return dec
def forward(self, input, sample_posterior=True):
if input.ndim==4:
input = input.unsqueeze(2)
posterior = self.encode(input)
if sample_posterior:
z = posterior.sample()
else:
z = posterior.mode()
dec = self.decode(z)
return dec, posterior
def get_input(self, batch, k):
x = batch[k]
if x.ndim==5:
x = x.permute(0, 4, 1, 2, 3).to(memory_format=torch.contiguous_format).float()
return x
if len(x.shape) == 3:
x = x[..., None]
x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format).float()
return x
def training_step(self, batch, batch_idx, optimizer_idx):
# tic = time.time()
inputs = self.get_input(batch, self.image_key)
# print(f"get_input time {time.time() - tic}")
# tic = time.time()
reconstructions, posterior = self(inputs)
# print(f"model forward time {time.time() - tic}")
if optimizer_idx == 0:
# train encoder+decoder+logvar
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("aeloss", aeloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=False)
# print(f"cal loss time {time.time() - tic}")
return aeloss
if optimizer_idx == 1:
# train the discriminator
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("discloss", discloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=False)
# print(f"cal loss time {time.time() - tic}")
return discloss
def validation_step(self, batch, batch_idx):
with torch.no_grad():
inputs = self.get_input(batch, self.image_key)
reconstructions, posterior = self(inputs)
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, 0, self.global_step,
last_layer=self.get_last_layer(), split="val")
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, 1, self.global_step,
last_layer=self.get_last_layer(), split="val")
self.log("val/rec_loss", log_dict_ae["val/rec_loss"])
self.log_dict(log_dict_ae)
self.log_dict(log_dict_disc)
return self.log_dict
def configure_optimizers(self):
lr = self.learning_rate
opt_ae = torch.optim.Adam(list(self.encoder.parameters())+
list(self.decoder.parameters())+
list(self.quant_conv.parameters())+
list(self.post_quant_conv.parameters()),
lr=lr, betas=(0.5, 0.9))
opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(),
lr=lr, betas=(0.5, 0.9))
return [opt_ae, opt_disc], []
def get_last_layer(self):
return self.decoder.conv_out.weight
@torch.no_grad()
def log_images(self, batch, only_inputs=False, **kwargs):
log = dict()
x = self.get_input(batch, self.image_key)
x = x.to(self.device)
if not only_inputs:
xrec, posterior = self(x)
if x.shape[1] > 3:
# colorize with random projection
assert xrec.shape[1] > 3
x = self.to_rgb(x)
xrec = self.to_rgb(xrec)
log["samples"] = self.decode(torch.randn_like(posterior.sample()))
log["reconstructions"] = xrec
log["inputs"] = x
return log
def to_rgb(self, x):
assert self.image_key == "segmentation"
if not hasattr(self, "colorize"):
self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x))
x = F.conv2d(x, weight=self.colorize)
x = 2.*(x-x.min())/(x.max()-x.min()) - 1.
return x
class AutoencoderKL(pl.LightningModule):
def __init__(self,
ddconfig,
lossconfig,
embed_dim,
ckpt_path=None,
ignore_keys=[],
image_key="image",
colorize_nlabels=None,
monitor=None,
):
super().__init__()
self.image_key = image_key
self.encoder = Encoder(**ddconfig)
self.decoder = Decoder(**ddconfig)
self.loss = instantiate_from_config(lossconfig)
assert ddconfig["double_z"]
self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1)
self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
self.embed_dim = embed_dim
if colorize_nlabels is not None:
assert type(colorize_nlabels)==int
self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1))
if monitor is not None:
self.monitor = monitor
if ckpt_path is not None:
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)
def init_from_ckpt(self, path, ignore_keys=list()):
sd = torch.load(path, map_location="cpu")["state_dict"]
keys = list(sd.keys())
for k in keys:
for ik in ignore_keys:
if k.startswith(ik):
print("Deleting key {} from state_dict.".format(k))
del sd[k]
self.load_state_dict(sd, strict=False)
print(f"Restored from {path}")
def encode(self, x):
h = self.encoder(x)
moments = self.quant_conv(h)
posterior = DiagonalGaussianDistribution(moments)
return posterior
def decode(self, z):
z = self.post_quant_conv(z)
dec = self.decoder(z)
return dec
def forward(self, input, sample_posterior=True):
posterior = self.encode(input)
if sample_posterior:
z = posterior.sample()
else:
z = posterior.mode()
dec = self.decode(z)
return dec, posterior
def get_input(self, batch, k):
x = batch[k]
if len(x.shape) == 3:
x = x[..., None]
x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format).float()
return x
def training_step(self, batch, batch_idx, optimizer_idx):
# tic = time.time()
inputs = self.get_input(batch, self.image_key)
# print(f"get_input time {time.time() - tic}")
# tic = time.time()
reconstructions, posterior = self(inputs)
# print(f"model forward time {time.time() - tic}")
tic = time.time()
if optimizer_idx == 0:
# train encoder+decoder+logvar
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("aeloss", aeloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=False)
# print(f"cal loss time {time.time() - tic}")
return aeloss
if optimizer_idx == 1:
# train the discriminator
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("discloss", discloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=False)
# print(f"cal loss time {time.time() - tic}")
return discloss
def validation_step(self, batch, batch_idx):
tic = time.time()
inputs = self.get_input(batch, self.image_key)
print(f"get_input time {time.time() - tic}")
tic = time.time()
reconstructions, posterior = self(inputs)
print(f"val forward time {time.time() - tic}")
tic = time.time()
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, 0, self.global_step,
last_layer=self.get_last_layer(), split="val")
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, 1, self.global_step,
last_layer=self.get_last_layer(), split="val")
self.log("val/rec_loss", log_dict_ae["val/rec_loss"])
self.log_dict(log_dict_ae)
self.log_dict(log_dict_disc)
print(f"val end time {time.time() - tic}")
return self.log_dict
def configure_optimizers(self):
lr = self.learning_rate
opt_ae = torch.optim.Adam(list(self.encoder.parameters())+
list(self.decoder.parameters())+
list(self.quant_conv.parameters())+
list(self.post_quant_conv.parameters()),
lr=lr, betas=(0.5, 0.9))
opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(),
lr=lr, betas=(0.5, 0.9))
return [opt_ae, opt_disc], []
def get_last_layer(self):
return self.decoder.conv_out.weight
@torch.no_grad()
def log_images(self, batch, only_inputs=False, **kwargs):
log = dict()
x = self.get_input(batch, self.image_key)
x = x.to(self.device)
if not only_inputs:
xrec, posterior = self(x)
if x.shape[1] > 3:
# colorize with random projection
assert xrec.shape[1] > 3
x = self.to_rgb(x)
xrec = self.to_rgb(xrec)
log["samples"] = self.decode(torch.randn_like(posterior.sample()))
log["reconstructions"] = xrec
log["inputs"] = x
return log
def to_rgb(self, x):
assert self.image_key == "segmentation"
if not hasattr(self, "colorize"):
self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x))
x = F.conv2d(x, weight=self.colorize)
x = 2.*(x-x.min())/(x.max()-x.min()) - 1.
return x
class IdentityFirstStage(torch.nn.Module):
def __init__(self, *args, vq_interface=False, **kwargs):
self.vq_interface = vq_interface # TODO: Should be true by default but check to not break older stuff
super().__init__()
def encode(self, x, *args, **kwargs):
return x
def decode(self, x, *args, **kwargs):
return x
def quantize(self, x, *args, **kwargs):
if self.vq_interface:
return x, None, [None, None, None]
return x
def forward(self, x, *args, **kwargs):
return x
@@ -0,0 +1,326 @@
from dataclasses import dataclass
from typing import Dict, Optional, Tuple
import numpy as np
import pytorch_lightning as pl
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from ..util import instantiate_from_config
from .cogvideox_enc_dec import (CogVideoXDecoder3D, CogVideoXEncoder3D,
CogVideoXSafeConv3d)
class DiagonalGaussianDistribution:
def __init__(
self,
mean: torch.Tensor,
logvar: torch.Tensor,
deterministic: bool = False,
):
self.mean = mean
self.logvar = torch.clamp(logvar, -30.0, 20.0)
self.deterministic = deterministic
if deterministic:
self.var = self.std = torch.zeros_like(self.mean)
else:
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
def sample(self, generator = None) -> torch.FloatTensor:
x = torch.randn(
self.mean.shape,
generator=generator,
device=self.mean.device,
dtype=self.mean.dtype,
)
return self.mean + self.std * x
def mode(self):
return self.mean
def kl(self, other: Optional["DiagonalGaussianDistribution"] = None) -> torch.Tensor:
dims = list(range(1, self.mean.ndim))
if self.deterministic:
return torch.Tensor([0.0])
else:
if other is None:
return 0.5 * torch.sum(
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar,
dim=dims,
)
else:
return 0.5 * torch.sum(
torch.pow(self.mean - other.mean, 2) / other.var
+ self.var / other.var
- 1.0
- self.logvar
+ other.logvar,
dim=dims,
)
def nll(self, sample: torch.Tensor) -> torch.Tensor:
dims = list(range(1, self.mean.ndim))
if self.deterministic:
return torch.Tensor([0.0])
logtwopi = np.log(2.0 * np.pi)
return 0.5 * torch.sum(
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
dim=dims,
)
@dataclass
class EncoderOutput:
latent_dist: DiagonalGaussianDistribution
@dataclass
class DecoderOutput:
sample: torch.Tensor
def str_eval(item):
if type(item) == str:
return eval(item)
else:
return item
class AutoencoderKLMagvit_CogVideoX(pl.LightningModule):
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: Tuple[str] = (
"CogVideoXDownBlock3D",
"CogVideoXDownBlock3D",
"CogVideoXDownBlock3D",
"CogVideoXDownBlock3D",
),
up_block_types: Tuple[str] = (
"CogVideoXUpBlock3D",
"CogVideoXUpBlock3D",
"CogVideoXUpBlock3D",
"CogVideoXUpBlock3D",
),
block_out_channels: Tuple[int] = (128, 256, 256, 512),
latent_channels: int = 16,
layers_per_block: int = 3,
act_fn: str = "silu",
norm_eps: float = 1e-6,
norm_num_groups: int = 32,
temporal_compression_ratio: float = 4,
use_quant_conv: bool = False,
use_post_quant_conv: bool = False,
mini_batch_encoder=4,
mini_batch_decoder=1,
image_key="image",
train_decoder_only=False,
train_encoder_only=False,
monitor=None,
ckpt_path=None,
lossconfig=None,
):
super().__init__()
self.image_key = image_key
down_block_types = str_eval(down_block_types)
up_block_types = str_eval(up_block_types)
self.encoder = CogVideoXEncoder3D(
in_channels=in_channels,
out_channels=latent_channels,
down_block_types=down_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
act_fn=act_fn,
norm_eps=norm_eps,
norm_num_groups=norm_num_groups,
temporal_compression_ratio=temporal_compression_ratio,
)
self.decoder = CogVideoXDecoder3D(
in_channels=latent_channels,
out_channels=out_channels,
up_block_types=up_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
act_fn=act_fn,
norm_eps=norm_eps,
norm_num_groups=norm_num_groups,
temporal_compression_ratio=temporal_compression_ratio,
)
self.quant_conv = CogVideoXSafeConv3d(2 * out_channels, 2 * out_channels, 1) if use_quant_conv else None
self.post_quant_conv = CogVideoXSafeConv3d(out_channels, out_channels, 1) if use_post_quant_conv else None
self.mini_batch_encoder = mini_batch_encoder
self.mini_batch_decoder = mini_batch_decoder
self.train_decoder_only = train_decoder_only
self.train_encoder_only = train_encoder_only
if train_decoder_only:
self.encoder.requires_grad_(False)
if self.quant_conv is not None:
self.quant_conv.requires_grad_(False)
if train_encoder_only:
self.decoder.requires_grad_(False)
if self.post_quant_conv is not None:
self.post_quant_conv.requires_grad_(False)
if monitor is not None:
self.monitor = monitor
if ckpt_path is not None:
self.init_from_ckpt(ckpt_path, ignore_keys="loss")
if lossconfig is not None:
self.loss = instantiate_from_config(lossconfig)
def init_from_ckpt(self, path, ignore_keys=list()):
if path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
sd = load_file(path)
else:
sd = torch.load(path, map_location="cpu")
if "state_dict" in list(sd.keys()):
sd = sd["state_dict"]
keys = list(sd.keys())
for k in keys:
for ik in ignore_keys:
if k.startswith(ik):
print("Deleting key {} from state_dict.".format(k))
del sd[k]
m, u = self.load_state_dict(sd, strict=False) # loss.item can be ignored successfully
print(f"Restored from {path}")
print(f"missing keys: {str(m)}, unexpected keys: {str(u)}")
def encode(self, x: torch.Tensor) -> EncoderOutput:
h = self.encoder(x)
self.encoder._clear_fake_context_parallel_cache()
if self.quant_conv is not None:
moments: torch.Tensor = self.quant_conv(h)
else:
moments: torch.Tensor = h
mean, logvar = moments.chunk(2, dim=1)
posterior = DiagonalGaussianDistribution(mean, logvar)
return posterior
def decode(self, z: torch.Tensor) -> DecoderOutput:
if self.post_quant_conv is not None:
z = self.post_quant_conv(z)
decoded = self.decoder(z)
self.decoder._clear_fake_context_parallel_cache()
return decoded
def forward(self, input, sample_posterior=True):
if input.ndim==4:
input = input.unsqueeze(2)
posterior = self.encode(input)
if sample_posterior:
z = posterior.sample()
else:
z = posterior.mode()
# print("stt latent shape", z.shape)
dec = self.decode(z)
return dec, posterior
def get_input(self, batch, k):
x = batch[k]
if x.ndim==5:
x = x.permute(0, 4, 1, 2, 3).to(memory_format=torch.contiguous_format).float()
return x
if len(x.shape) == 3:
x = x[..., None]
x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format).float()
return x
def training_step(self, batch, batch_idx, optimizer_idx):
inputs = self.get_input(batch, self.image_key)
reconstructions, posterior = self(inputs)
if optimizer_idx == 0:
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("aeloss", aeloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=False)
return aeloss
if optimizer_idx == 1:
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("discloss", discloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=False)
return discloss
def validation_step(self, batch, batch_idx):
with torch.no_grad():
inputs = self.get_input(batch, self.image_key)
reconstructions, posterior = self(inputs)
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, 0, self.global_step,
last_layer=self.get_last_layer(), split="val")
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, 1, self.global_step,
last_layer=self.get_last_layer(), split="val")
self.log("val/rec_loss", log_dict_ae["val/rec_loss"])
self.log_dict(log_dict_ae)
self.log_dict(log_dict_disc)
return self.log_dict
def configure_optimizers(self):
lr = self.learning_rate
if self.train_decoder_only:
if self.post_quant_conv is not None:
training_list = list(self.decoder.parameters()) + list(self.post_quant_conv.parameters())
else:
training_list = list(self.decoder.parameters())
opt_ae = torch.optim.Adam(training_list, lr=lr, betas=(0.5, 0.9))
elif self.train_encoder_only:
if self.quant_conv is not None:
training_list = list(self.encoder.parameters()) + list(self.quant_conv.parameters())
else:
training_list = list(self.encoder.parameters())
opt_ae = torch.optim.Adam(training_list, lr=lr, betas=(0.5, 0.9))
else:
training_list = list(self.encoder.parameters()) + list(self.decoder.parameters())
if self.quant_conv is not None:
training_list = training_list + list(self.quant_conv.parameters())
if self.post_quant_conv is not None:
training_list = training_list + list(self.post_quant_conv.parameters())
opt_ae = torch.optim.Adam(training_list, lr=lr, betas=(0.5, 0.9))
opt_disc = torch.optim.Adam(
list(self.loss.discriminator3d.parameters()) + list(self.loss.discriminator.parameters()),
lr=lr, betas=(0.5, 0.9)
)
return [opt_ae, opt_disc], []
def get_last_layer(self):
return self.decoder.conv_out.conv.weight
@torch.no_grad()
def log_images(self, batch, only_inputs=False, **kwargs):
log = dict()
x = self.get_input(batch, self.image_key)
x = x.to(self.device)
if not only_inputs:
xrec, posterior = self(x)
if x.shape[1] > 3:
# colorize with random projection
assert xrec.shape[1] > 3
x = self.to_rgb(x)
xrec = self.to_rgb(xrec)
log["samples"] = self.decode(torch.randn_like(posterior.sample()))
log["reconstructions"] = xrec
log["inputs"] = x
return log
def to_rgb(self, x):
assert self.image_key == "segmentation"
if not hasattr(self, "colorize"):
self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x))
x = F.conv2d(x, weight=self.colorize)
x = 2.*(x-x.min())/(x.max()-x.min()) - 1.
return x
@@ -0,0 +1,312 @@
# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
# All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Optional, Tuple
import numpy as np
import torch
import torch.nn as nn
from diffusers.models.autoencoders.autoencoder_kl_cogvideox import (
CogVideoXCausalConv3d, CogVideoXDownBlock3D, CogVideoXMidBlock3D,
CogVideoXSafeConv3d, CogVideoXSpatialNorm3D, CogVideoXUpBlock3D)
from diffusers.utils import logging
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
class CogVideoXEncoder3D(nn.Module):
r"""
The `CogVideoXEncoder3D` layer of a variational autoencoder that encodes its input into a latent representation.
Args:
in_channels (`int`, *optional*, defaults to 3):
The number of input channels.
out_channels (`int`, *optional*, defaults to 3):
The number of output channels.
down_block_types (`Tuple[str, ...]`, *optional*, defaults to `("DownEncoderBlock2D",)`):
The types of down blocks to use. See `~diffusers.models.unet_2d_blocks.get_down_block` for available
options.
block_out_channels (`Tuple[int, ...]`, *optional*, defaults to `(64,)`):
The number of output channels for each block.
act_fn (`str`, *optional*, defaults to `"silu"`):
The activation function to use. See `~diffusers.models.activations.get_activation` for available options.
layers_per_block (`int`, *optional*, defaults to 2):
The number of layers per block.
norm_num_groups (`int`, *optional*, defaults to 32):
The number of groups for normalization.
"""
_supports_gradient_checkpointing = True
def __init__(
self,
in_channels: int = 3,
out_channels: int = 16,
down_block_types: Tuple[str, ...] = (
"CogVideoXDownBlock3D",
"CogVideoXDownBlock3D",
"CogVideoXDownBlock3D",
"CogVideoXDownBlock3D",
),
block_out_channels: Tuple[int, ...] = (128, 256, 256, 512),
layers_per_block: int = 3,
act_fn: str = "silu",
norm_eps: float = 1e-6,
norm_num_groups: int = 32,
dropout: float = 0.0,
pad_mode: str = "first",
temporal_compression_ratio: float = 4,
):
super().__init__()
# log2 of temporal_compress_times
temporal_compress_level = int(np.log2(temporal_compression_ratio))
self.conv_in = CogVideoXCausalConv3d(in_channels, block_out_channels[0], kernel_size=3, pad_mode=pad_mode)
self.down_blocks = nn.ModuleList([])
# down blocks
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
compress_time = i < temporal_compress_level
if down_block_type == "CogVideoXDownBlock3D":
down_block = CogVideoXDownBlock3D(
in_channels=input_channel,
out_channels=output_channel,
temb_channels=0,
dropout=dropout,
num_layers=layers_per_block,
resnet_eps=norm_eps,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
add_downsample=not is_final_block,
compress_time=compress_time,
)
else:
raise ValueError("Invalid `down_block_type` encountered. Must be `CogVideoXDownBlock3D`")
self.down_blocks.append(down_block)
# mid block
self.mid_block = CogVideoXMidBlock3D(
in_channels=block_out_channels[-1],
temb_channels=0,
dropout=dropout,
num_layers=2,
resnet_eps=norm_eps,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
pad_mode=pad_mode,
)
self.norm_out = nn.GroupNorm(norm_num_groups, block_out_channels[-1], eps=1e-6)
self.conv_act = nn.SiLU()
self.conv_out = CogVideoXCausalConv3d(
block_out_channels[-1], 2 * out_channels, kernel_size=3, pad_mode=pad_mode
)
self.gradient_checkpointing = False
def _clear_fake_context_parallel_cache(self):
for name, module in self.named_modules():
if isinstance(module, CogVideoXCausalConv3d):
logger.debug(f"Clearing fake Context Parallel cache for layer: {name}")
module._clear_fake_context_parallel_cache()
def forward(self, sample: torch.Tensor, temb: Optional[torch.Tensor] = None) -> torch.Tensor:
r"""The forward method of the `CogVideoXEncoder3D` class."""
hidden_states = self.conv_in(sample)
if self.training and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
# 1. Down
for down_block in self.down_blocks:
hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(down_block), hidden_states, temb, None
)
# 2. Mid
hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.mid_block), hidden_states, temb, None
)
else:
# 1. Down
for down_block in self.down_blocks:
hidden_states = down_block(hidden_states, temb, None)
# 2. Mid
hidden_states = self.mid_block(hidden_states, temb, None)
# 3. Post-process
hidden_states = self.norm_out(hidden_states)
hidden_states = self.conv_act(hidden_states)
hidden_states = self.conv_out(hidden_states)
return hidden_states
class CogVideoXDecoder3D(nn.Module):
r"""
The `CogVideoXDecoder3D` layer of a variational autoencoder that decodes its latent representation into an output
sample.
Args:
in_channels (`int`, *optional*, defaults to 3):
The number of input channels.
out_channels (`int`, *optional*, defaults to 3):
The number of output channels.
up_block_types (`Tuple[str, ...]`, *optional*, defaults to `("UpDecoderBlock2D",)`):
The types of up blocks to use. See `~diffusers.models.unet_2d_blocks.get_up_block` for available options.
block_out_channels (`Tuple[int, ...]`, *optional*, defaults to `(64,)`):
The number of output channels for each block.
act_fn (`str`, *optional*, defaults to `"silu"`):
The activation function to use. See `~diffusers.models.activations.get_activation` for available options.
layers_per_block (`int`, *optional*, defaults to 2):
The number of layers per block.
norm_num_groups (`int`, *optional*, defaults to 32):
The number of groups for normalization.
"""
_supports_gradient_checkpointing = True
def __init__(
self,
in_channels: int = 16,
out_channels: int = 3,
up_block_types: Tuple[str, ...] = (
"CogVideoXUpBlock3D",
"CogVideoXUpBlock3D",
"CogVideoXUpBlock3D",
"CogVideoXUpBlock3D",
),
block_out_channels: Tuple[int, ...] = (128, 256, 256, 512),
layers_per_block: int = 3,
act_fn: str = "silu",
norm_eps: float = 1e-6,
norm_num_groups: int = 32,
dropout: float = 0.0,
pad_mode: str = "first",
temporal_compression_ratio: float = 4,
):
super().__init__()
reversed_block_out_channels = list(reversed(block_out_channels))
self.conv_in = CogVideoXCausalConv3d(
in_channels, reversed_block_out_channels[0], kernel_size=3, pad_mode=pad_mode
)
# mid block
self.mid_block = CogVideoXMidBlock3D(
in_channels=reversed_block_out_channels[0],
temb_channels=0,
num_layers=2,
resnet_eps=norm_eps,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
spatial_norm_dim=in_channels,
pad_mode=pad_mode,
)
# up blocks
self.up_blocks = nn.ModuleList([])
output_channel = reversed_block_out_channels[0]
temporal_compress_level = int(np.log2(temporal_compression_ratio))
for i, up_block_type in enumerate(up_block_types):
prev_output_channel = output_channel
output_channel = reversed_block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
compress_time = i < temporal_compress_level
if up_block_type == "CogVideoXUpBlock3D":
up_block = CogVideoXUpBlock3D(
in_channels=prev_output_channel,
out_channels=output_channel,
temb_channels=0,
dropout=dropout,
num_layers=layers_per_block + 1,
resnet_eps=norm_eps,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
spatial_norm_dim=in_channels,
add_upsample=not is_final_block,
compress_time=compress_time,
pad_mode=pad_mode,
)
prev_output_channel = output_channel
else:
raise ValueError("Invalid `up_block_type` encountered. Must be `CogVideoXUpBlock3D`")
self.up_blocks.append(up_block)
self.norm_out = CogVideoXSpatialNorm3D(reversed_block_out_channels[-1], in_channels, groups=norm_num_groups)
self.conv_act = nn.SiLU()
self.conv_out = CogVideoXCausalConv3d(
reversed_block_out_channels[-1], out_channels, kernel_size=3, pad_mode=pad_mode
)
self.gradient_checkpointing = False
def _clear_fake_context_parallel_cache(self):
for name, module in self.named_modules():
if isinstance(module, CogVideoXCausalConv3d):
logger.debug(f"Clearing fake Context Parallel cache for layer: {name}")
module._clear_fake_context_parallel_cache()
def forward(self, sample: torch.Tensor, temb: Optional[torch.Tensor] = None) -> torch.Tensor:
r"""The forward method of the `CogVideoXDecoder3D` class."""
hidden_states = self.conv_in(sample)
if self.training and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
# 1. Mid
hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.mid_block), hidden_states, temb, sample
)
# 2. Up
for up_block in self.up_blocks:
hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(up_block), hidden_states, temb, sample
)
else:
# 1. Mid
hidden_states = self.mid_block(hidden_states, temb, sample)
# 2. Up
for up_block in self.up_blocks:
hidden_states = up_block(hidden_states, temb, sample)
# 3. Post-process
hidden_states = self.norm_out(hidden_states, sample)
hidden_states = self.conv_act(hidden_states)
hidden_states = self.conv_out(hidden_states)
return hidden_states
+234
View File
@@ -0,0 +1,234 @@
import torch
import torch.nn.functional as F
from torch import nn
def cast_tuple(t, length = 1):
return t if isinstance(t, tuple) else ((t,) * length)
def divisible_by(num, den):
return (num % den) == 0
def is_odd(n):
return not divisible_by(n, 2)
class CausalConv3d(nn.Module):
def __init__(
self,
chan_in,
chan_out,
kernel_size,
pad_mode = 'constant',
**kwargs
):
super().__init__()
kernel_size = cast_tuple(kernel_size, 3)
time_kernel_size, height_kernel_size, width_kernel_size = kernel_size
assert is_odd(height_kernel_size) and is_odd(width_kernel_size)
dilation = kwargs.pop('dilation', 1)
stride = kwargs.pop('stride', 1)
self.pad_mode = pad_mode
time_pad = dilation * (time_kernel_size - 1) + (1 - stride)
height_pad = height_kernel_size // 2
width_pad = width_kernel_size // 2
self.time_pad = time_pad
self.time_causal_padding = (width_pad, width_pad, height_pad, height_pad, time_pad, 0)
stride = (stride, 1, 1)
dilation = (dilation, 1, 1)
self.conv = nn.Conv3d(chan_in, chan_out, kernel_size, stride = stride, dilation = dilation, **kwargs)
def forward(self, x):
x = F.pad(x, self.time_causal_padding, mode = 'replicate')
return self.conv(x)
class Swish(nn.Module):
def __init__(self) -> None:
super().__init__()
def forward(self, x):
return x * F.sigmoid(x)
class ResBlockX(nn.Module):
def __init__(self, inchannel) -> None:
super().__init__()
self.conv = nn.Sequential(
nn.GroupNorm(32, inchannel),
Swish(),
CausalConv3d(inchannel, inchannel, 3),
nn.GroupNorm(32, inchannel),
Swish(),
CausalConv3d(inchannel, inchannel, 3)
)
def forward(self, x):
return x + self.conv(x)
class ResBlockXY(nn.Module):
def __init__(self, inchannel, outchannel) -> None:
super().__init__()
self.conv = nn.Sequential(
nn.GroupNorm(32, inchannel),
Swish(),
CausalConv3d(inchannel, outchannel, 3),
nn.GroupNorm(32, outchannel),
Swish(),
CausalConv3d(outchannel, outchannel, 3)
)
self.conv_1 = nn.Conv3d(inchannel, outchannel, 1)
def forward(self, x):
return self.conv_1(x) + self.conv(x)
class PoolDown222(nn.Module):
def __init__(self) -> None:
super().__init__()
self.pool = nn.AvgPool3d(2, 2)
def forward(self, x):
x = F.pad(x, (0, 0, 0, 0, 1, 0), 'replicate')
return self.pool(x)
class PoolDown122(nn.Module):
def __init__(self) -> None:
super().__init__()
self.pool = nn.AvgPool3d((1, 2, 2), (1, 2, 2))
def forward(self, x):
return self.pool(x)
class Unpool222(nn.Module):
def __init__(self) -> None:
super().__init__()
self.up = nn.Upsample(scale_factor=2, mode='nearest')
def forward(self, x):
x = self.up(x)
return x[:, :, 1:]
class Unpool122(nn.Module):
def __init__(self) -> None:
super().__init__()
self.up = nn.Upsample(scale_factor=(1, 2, 2), mode='nearest')
def forward(self, x):
x = self.up(x)
return x
class ResBlockDown(nn.Module):
def __init__(self, inchannel, outchannel) -> None:
super().__init__()
self.blcok = nn.Sequential(
CausalConv3d(inchannel, outchannel, 3),
nn.LeakyReLU(inplace=True),
PoolDown222(),
CausalConv3d(outchannel, outchannel, 3),
nn.LeakyReLU(inplace=True)
)
self.res = nn.Sequential(
PoolDown222(),
nn.Conv3d(inchannel, outchannel, 1)
)
def forward(self, x):
return self.res(x) + self.blcok(x)
class Discriminator(nn.Module):
def __init__(self) -> None:
super().__init__()
self.block = nn.Sequential(
CausalConv3d(3, 64, 3),
nn.LeakyReLU(inplace=True),
ResBlockDown(64, 128),
ResBlockDown(128, 256),
ResBlockDown(256, 256),
ResBlockDown(256, 256),
ResBlockDown(256, 256),
CausalConv3d(256, 256, 3),
nn.LeakyReLU(inplace=True),
nn.AdaptiveAvgPool3d(1),
nn.Flatten(),
nn.Linear(256, 256),
nn.LeakyReLU(inplace=True),
nn.Linear(256, 1)
)
def forward(self, x):
if x.ndim==4:
x = x.unsqueeze(2)
return self.block(x)
class Encoder(nn.Module):
def __init__(self) -> None:
super().__init__()
self.encoder = nn.Sequential(
CausalConv3d(3, 64, 3),
ResBlockX(64),
ResBlockX(64),
PoolDown222(),
ResBlockXY(64, 128),
ResBlockX(128),
PoolDown222(),
ResBlockX(128),
ResBlockX(128),
PoolDown122(),
ResBlockXY(128, 256),
ResBlockX(256),
ResBlockX(256),
ResBlockX(256),
nn.GroupNorm(32, 256),
Swish(),
nn.Conv3d(256, 16, 1)
)
def forward(self, x):
return self.encoder(x)
class Decoder(nn.Module):
def __init__(self) -> None:
super().__init__()
self.decoder = nn.Sequential(
CausalConv3d(8, 256, 3),
ResBlockX(256),
ResBlockX(256),
ResBlockX(256),
ResBlockX(256),
Unpool122(),
CausalConv3d(256, 256, 3),
ResBlockXY(256, 128),
ResBlockX(128),
Unpool222(),
CausalConv3d(128, 128, 3),
ResBlockX(128),
ResBlockX(128),
Unpool222(),
CausalConv3d(128, 128, 3),
ResBlockXY(128, 64),
ResBlockX(64),
nn.GroupNorm(32, 64),
Swish(),
CausalConv3d(64, 64, 3)
)
self.conv_out = nn.Conv3d(64, 3, 1)
def forward(self, x):
return self.conv_out(self.decoder(x))
if __name__=='__main__':
encoder = Encoder()
decoder = Decoder()
dis = Discriminator()
x = torch.randn((1, 3, 1, 64, 64))
embedding = encoder(x)
y = decoder(embedding)
tmp = torch.randn((1, 4, 1, 64, 64))
print('something mmm')
@@ -0,0 +1,345 @@
from dataclasses import dataclass
from typing import Optional
import numpy as np
import pytorch_lightning as pl
import torch
import torch.nn as nn
import torch.nn.functional as F
from ..util import instantiate_from_config
from .omnigen_enc_dec import Decoder as omnigen_Mag_Decoder
from .omnigen_enc_dec import Encoder as omnigen_Mag_Encoder
class DiagonalGaussianDistribution:
def __init__(
self,
mean: torch.Tensor,
logvar: torch.Tensor,
deterministic: bool = False,
):
self.mean = mean
self.logvar = torch.clamp(logvar, -30.0, 20.0)
self.deterministic = deterministic
if deterministic:
self.var = self.std = torch.zeros_like(self.mean)
else:
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
def sample(self, generator = None) -> torch.FloatTensor:
x = torch.randn(
self.mean.shape,
generator=generator,
device=self.mean.device,
dtype=self.mean.dtype,
)
return self.mean + self.std * x
def mode(self):
return self.mean
def kl(self, other: Optional["DiagonalGaussianDistribution"] = None) -> torch.Tensor:
dims = list(range(1, self.mean.ndim))
if self.deterministic:
return torch.Tensor([0.0])
else:
if other is None:
return 0.5 * torch.sum(
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar,
dim=dims,
)
else:
return 0.5 * torch.sum(
torch.pow(self.mean - other.mean, 2) / other.var
+ self.var / other.var
- 1.0
- self.logvar
+ other.logvar,
dim=dims,
)
def nll(self, sample: torch.Tensor) -> torch.Tensor:
dims = list(range(1, self.mean.ndim))
if self.deterministic:
return torch.Tensor([0.0])
logtwopi = np.log(2.0 * np.pi)
return 0.5 * torch.sum(
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
dim=dims,
)
@dataclass
class EncoderOutput:
latent_dist: DiagonalGaussianDistribution
@dataclass
class DecoderOutput:
sample: torch.Tensor
def str_eval(item):
if type(item) == str:
return eval(item)
else:
return item
class AutoencoderKLMagvit_fromOmnigen(pl.LightningModule):
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
ch = 128,
ch_mult = [ 1,2,4,4 ],
block_out_channels = [128, 256, 512, 512],
use_gc_blocks = None,
down_block_types: tuple = None,
up_block_types: tuple = None,
mid_block_type: str = "MidBlock3D",
mid_block_use_attention: bool = True,
mid_block_attention_type: str = "3d",
mid_block_num_attention_heads: int = 1,
layers_per_block: int = 2,
act_fn: str = "silu",
num_attention_heads: int = 1,
latent_channels: int = 4,
norm_num_groups: int = 32,
image_key="image",
monitor=None,
ckpt_path=None,
lossconfig=None,
slice_mag_vae=False,
slice_compression_vae=False,
cache_compression_vae=False,
cache_mag_vae=False,
spatial_group_norm=False,
mini_batch_encoder=9,
mini_batch_decoder=3,
train_decoder_only=False,
train_encoder_only=False,
):
super().__init__()
self.image_key = image_key
down_block_types = str_eval(down_block_types)
up_block_types = str_eval(up_block_types)
self.encoder = omnigen_Mag_Encoder(
in_channels=in_channels,
out_channels=latent_channels,
down_block_types=down_block_types,
ch=ch,
ch_mult=ch_mult,
block_out_channels=block_out_channels,
use_gc_blocks=use_gc_blocks,
mid_block_type=mid_block_type,
mid_block_use_attention=mid_block_use_attention,
mid_block_attention_type=mid_block_attention_type,
mid_block_num_attention_heads=mid_block_num_attention_heads,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
num_attention_heads=num_attention_heads,
double_z=True,
slice_mag_vae=slice_mag_vae,
slice_compression_vae=slice_compression_vae,
cache_compression_vae=cache_compression_vae,
cache_mag_vae=cache_mag_vae,
spatial_group_norm=spatial_group_norm,
mini_batch_encoder=mini_batch_encoder,
)
self.decoder = omnigen_Mag_Decoder(
in_channels=latent_channels,
out_channels=out_channels,
up_block_types=up_block_types,
ch=ch,
ch_mult=ch_mult,
block_out_channels=block_out_channels,
use_gc_blocks=use_gc_blocks,
mid_block_type=mid_block_type,
mid_block_use_attention=mid_block_use_attention,
mid_block_attention_type=mid_block_attention_type,
mid_block_num_attention_heads=mid_block_num_attention_heads,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
num_attention_heads=num_attention_heads,
slice_mag_vae=slice_mag_vae,
slice_compression_vae=slice_compression_vae,
cache_compression_vae=cache_compression_vae,
cache_mag_vae=cache_mag_vae,
spatial_group_norm=spatial_group_norm,
mini_batch_decoder=mini_batch_decoder,
)
self.quant_conv = nn.Conv3d(2 * latent_channels, 2 * latent_channels, kernel_size=1)
self.post_quant_conv = nn.Conv3d(latent_channels, latent_channels, kernel_size=1)
self.mini_batch_encoder = mini_batch_encoder
self.mini_batch_decoder = mini_batch_decoder
self.train_decoder_only = train_decoder_only
self.train_encoder_only = train_encoder_only
if train_decoder_only:
self.encoder.requires_grad_(False)
if self.quant_conv is not None:
self.quant_conv.requires_grad_(False)
if train_encoder_only:
self.decoder.requires_grad_(False)
if self.post_quant_conv is not None:
self.post_quant_conv.requires_grad_(False)
if monitor is not None:
self.monitor = monitor
if ckpt_path is not None:
self.init_from_ckpt(ckpt_path, ignore_keys="loss")
if lossconfig is not None:
self.loss = instantiate_from_config(lossconfig)
def init_from_ckpt(self, path, ignore_keys=list()):
if path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
sd = load_file(path)
else:
sd = torch.load(path, map_location="cpu")
if "state_dict" in list(sd.keys()):
sd = sd["state_dict"]
keys = list(sd.keys())
for k in keys:
for ik in ignore_keys:
if k.startswith(ik):
print("Deleting key {} from state_dict.".format(k))
del sd[k]
m, u = self.load_state_dict(sd, strict=False) # loss.item can be ignored successfully
print(f"Restored from {path}")
print(f"missing keys: {str(m)}, unexpected keys: {str(u)}")
def encode(self, x: torch.Tensor) -> EncoderOutput:
h = self.encoder(x)
if self.quant_conv is not None:
moments: torch.Tensor = self.quant_conv(h)
else:
moments: torch.Tensor = h
mean, logvar = moments.chunk(2, dim=1)
posterior = DiagonalGaussianDistribution(mean, logvar)
return posterior
def decode(self, z: torch.Tensor) -> DecoderOutput:
if self.post_quant_conv is not None:
z = self.post_quant_conv(z)
decoded = self.decoder(z)
return decoded
def forward(self, input, sample_posterior=True):
if input.ndim==4:
input = input.unsqueeze(2)
posterior = self.encode(input)
if sample_posterior:
z = posterior.sample()
else:
z = posterior.mode()
# print("stt latent shape", z.shape)
dec = self.decode(z)
return dec, posterior
def get_input(self, batch, k):
x = batch[k]
if x.ndim==5:
x = x.permute(0, 4, 1, 2, 3).to(memory_format=torch.contiguous_format).float()
return x
if len(x.shape) == 3:
x = x[..., None]
x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format).float()
return x
def training_step(self, batch, batch_idx, optimizer_idx):
inputs = self.get_input(batch, self.image_key)
reconstructions, posterior = self(inputs)
if optimizer_idx == 0:
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("aeloss", aeloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=False)
return aeloss
if optimizer_idx == 1:
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train")
self.log("discloss", discloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=False)
return discloss
def validation_step(self, batch, batch_idx):
with torch.no_grad():
inputs = self.get_input(batch, self.image_key)
reconstructions, posterior = self(inputs)
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, 0, self.global_step,
last_layer=self.get_last_layer(), split="val")
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, 1, self.global_step,
last_layer=self.get_last_layer(), split="val")
self.log("val/rec_loss", log_dict_ae["val/rec_loss"])
self.log_dict(log_dict_ae)
self.log_dict(log_dict_disc)
return self.log_dict
def configure_optimizers(self):
lr = self.learning_rate
if self.train_decoder_only:
if self.post_quant_conv is not None:
training_list = list(self.decoder.parameters()) + list(self.post_quant_conv.parameters())
else:
training_list = list(self.decoder.parameters())
opt_ae = torch.optim.Adam(training_list, lr=lr, betas=(0.5, 0.9))
elif self.train_encoder_only:
if self.quant_conv is not None:
training_list = list(self.encoder.parameters()) + list(self.quant_conv.parameters())
else:
training_list = list(self.encoder.parameters())
opt_ae = torch.optim.Adam(training_list, lr=lr, betas=(0.5, 0.9))
else:
training_list = list(self.encoder.parameters()) + list(self.decoder.parameters())
if self.quant_conv is not None:
training_list = training_list + list(self.quant_conv.parameters())
if self.post_quant_conv is not None:
training_list = training_list + list(self.post_quant_conv.parameters())
opt_ae = torch.optim.Adam(training_list, lr=lr, betas=(0.5, 0.9))
opt_disc = torch.optim.Adam(
list(self.loss.discriminator3d.parameters()) + list(self.loss.discriminator.parameters()),
lr=lr, betas=(0.5, 0.9)
)
return [opt_ae, opt_disc], []
def get_last_layer(self):
return self.decoder.conv_out.weight
@torch.no_grad()
def log_images(self, batch, only_inputs=False, **kwargs):
log = dict()
x = self.get_input(batch, self.image_key)
x = x.to(self.device)
if not only_inputs:
xrec, posterior = self(x)
if x.shape[1] > 3:
# colorize with random projection
assert xrec.shape[1] > 3
x = self.to_rgb(x)
xrec = self.to_rgb(xrec)
log["samples"] = self.decode(torch.randn_like(posterior.sample()))
log["reconstructions"] = xrec
log["inputs"] = x
return log
def to_rgb(self, x):
assert self.image_key == "segmentation"
if not hasattr(self, "colorize"):
self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x))
x = F.conv2d(x, weight=self.colorize)
x = 2.*(x-x.min())/(x.max()-x.min()) - 1.
return x
@@ -0,0 +1,669 @@
from typing import Any, Dict
import torch
import torch.nn as nn
from diffusers.utils import is_torch_version
from einops import rearrange
from ..modules.vaemodules.activations import get_activation
from ..modules.vaemodules.common import CausalConv3d
from ..modules.vaemodules.down_blocks import get_down_block
from ..modules.vaemodules.mid_blocks import get_mid_block
from ..modules.vaemodules.up_blocks import get_up_block
def create_custom_forward(module, return_dict=None):
def custom_forward(*inputs):
if return_dict is not None:
return module(*inputs, return_dict=return_dict)
else:
return module(*inputs)
return custom_forward
class Encoder(nn.Module):
r"""
The `Encoder` layer of a variational autoencoder that encodes its input into a latent representation.
Args:
in_channels (`int`, *optional*, defaults to 3):
The number of input channels.
out_channels (`int`, *optional*, defaults to 8):
The number of output channels.
down_block_types (`Tuple[str, ...]`, *optional*, defaults to `("SpatialDownBlock3D",)`):
The types of down blocks to use.
block_out_channels (`Tuple[int, ...]`, *optional*, defaults to `(64,)`):
The number of output channels for each block.
use_gc_blocks (`Tuple[bool, ...]`, *optional*, defaults to `None`):
Whether to use global context blocks for each down block.
mid_block_type (`str`, *optional*, defaults to `"MidBlock3D"`):
The type of mid block to use.
layers_per_block (`int`, *optional*, defaults to 2):
The number of layers per block.
norm_num_groups (`int`, *optional*, defaults to 32):
The number of groups for normalization.
act_fn (`str`, *optional*, defaults to `"silu"`):
The activation function to use. See `~diffusers.models.activations.get_activation` for available options.
num_attention_heads (`int`, *optional*, defaults to 1):
The number of attention heads to use.
double_z (`bool`, *optional*, defaults to `True`):
Whether to double the number of output channels for the last block.
"""
def __init__(
self,
in_channels: int = 3,
out_channels: int = 8,
down_block_types = ("SpatialDownBlock3D",),
ch = 128,
ch_mult = [1,2,4,4,],
block_out_channels = [128, 256, 512, 512],
use_gc_blocks = None,
mid_block_type: str = "MidBlock3D",
mid_block_use_attention: bool = True,
mid_block_attention_type: str = "3d",
mid_block_num_attention_heads: int = 1,
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
num_attention_heads: int = 1,
double_z: bool = True,
slice_mag_vae: bool = False,
slice_compression_vae: bool = False,
cache_compression_vae: bool = False,
cache_mag_vae: bool = False,
spatial_group_norm: bool = False,
mini_batch_encoder: int = 9,
verbose = False,
):
super().__init__()
if block_out_channels is None:
block_out_channels = [ch * i for i in ch_mult]
assert len(down_block_types) == len(block_out_channels), (
"Number of down block types must match number of block output channels."
)
if use_gc_blocks is not None:
assert len(use_gc_blocks) == len(down_block_types), (
"Number of GC blocks must match number of down block types."
)
else:
use_gc_blocks = [False] * len(down_block_types)
self.conv_in = CausalConv3d(
in_channels,
block_out_channels[0],
kernel_size=3,
)
self.down_blocks = nn.ModuleList([])
output_channels = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channels = output_channels
output_channels = block_out_channels[i]
is_final_block = (i == len(block_out_channels) - 1)
down_block = get_down_block(
down_block_type,
in_channels=input_channels,
out_channels=output_channels,
num_layers=layers_per_block,
act_fn=act_fn,
norm_num_groups=norm_num_groups,
norm_eps=1e-6,
num_attention_heads=num_attention_heads,
add_gc_block=use_gc_blocks[i],
add_downsample=not is_final_block,
)
self.down_blocks.append(down_block)
self.mid_block = get_mid_block(
mid_block_type,
in_channels=block_out_channels[-1],
num_layers=layers_per_block,
act_fn=act_fn,
norm_num_groups=norm_num_groups,
norm_eps=1e-6,
add_attention=mid_block_use_attention,
attention_type=mid_block_attention_type,
num_attention_heads=mid_block_num_attention_heads,
)
self.conv_norm_out = nn.GroupNorm(
num_channels=block_out_channels[-1],
num_groups=norm_num_groups,
eps=1e-6,
)
self.conv_act = get_activation(act_fn)
conv_out_channels = 2 * out_channels if double_z else out_channels
self.conv_out = CausalConv3d(block_out_channels[-1], conv_out_channels, kernel_size=3)
self.slice_mag_vae = slice_mag_vae
self.slice_compression_vae = slice_compression_vae
self.cache_compression_vae = cache_compression_vae
self.cache_mag_vae = cache_mag_vae
self.mini_batch_encoder = mini_batch_encoder
self.spatial_group_norm = spatial_group_norm
self.verbose = verbose
def set_padding_one_frame(self):
def _set_padding_one_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 1
for sub_name, sub_mod in module.named_children():
_set_padding_one_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_padding_one_frame(name, module)
def set_padding_more_frame(self):
def _set_padding_more_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 2
for sub_name, sub_mod in module.named_children():
_set_padding_more_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_padding_more_frame(name, module)
def set_magvit_padding_one_frame(self):
def _set_magvit_padding_one_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 3
for sub_name, sub_mod in module.named_children():
_set_magvit_padding_one_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_magvit_padding_one_frame(name, module)
def set_magvit_padding_more_frame(self):
def _set_magvit_padding_more_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 4
for sub_name, sub_mod in module.named_children():
_set_magvit_padding_more_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_magvit_padding_more_frame(name, module)
def set_cache_slice_vae_padding_one_frame(self):
def _set_cache_slice_vae_padding_one_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 5
for sub_name, sub_mod in module.named_children():
_set_cache_slice_vae_padding_one_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_cache_slice_vae_padding_one_frame(name, module)
def set_cache_slice_vae_padding_more_frame(self):
def _set_cache_slice_vae_padding_more_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 6
for sub_name, sub_mod in module.named_children():
_set_cache_slice_vae_padding_more_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_cache_slice_vae_padding_more_frame(name, module)
def set_3dgroupnorm_for_submodule(self):
def _set_3dgroupnorm_for_submodule(name, module):
if hasattr(module, 'set_3dgroupnorm'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.set_3dgroupnorm = True
for sub_name, sub_mod in module.named_children():
_set_3dgroupnorm_for_submodule(sub_name, sub_mod)
for name, module in self.named_children():
_set_3dgroupnorm_for_submodule(name, module)
def single_forward(self, x: torch.Tensor, previous_features: torch.Tensor, after_features: torch.Tensor) -> torch.Tensor:
# x: (B, C, T, H, W)
if self.training:
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
if previous_features is not None and after_features is None:
x = torch.concat([previous_features, x], 2)
elif previous_features is None and after_features is not None:
x = torch.concat([x, after_features], 2)
elif previous_features is not None and after_features is not None:
x = torch.concat([previous_features, x, after_features], 2)
if self.training:
x = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.conv_in),
x,
**ckpt_kwargs,
)
else:
x = self.conv_in(x)
for down_block in self.down_blocks:
if self.training:
x = torch.utils.checkpoint.checkpoint(
create_custom_forward(down_block),
x,
**ckpt_kwargs,
)
else:
x = down_block(x)
x = self.mid_block(x)
if self.spatial_group_norm:
batch_size = x.shape[0]
x = rearrange(x, "b c t h w -> (b t) c h w")
x = self.conv_norm_out(x)
x = rearrange(x, "(b t) c h w -> b c t h w", b=batch_size)
else:
x = self.conv_norm_out(x)
x = self.conv_act(x)
x = self.conv_out(x)
if previous_features is not None and after_features is None:
x = x[:, :, 1:]
elif previous_features is None and after_features is not None:
x = x[:, :, :2]
elif previous_features is not None and after_features is not None:
x = x[:, :, 1:3]
return x
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.spatial_group_norm:
self.set_3dgroupnorm_for_submodule()
if self.cache_mag_vae:
self.set_magvit_padding_one_frame()
first_frames = self.single_forward(x[:, :, 0:1, :, :], None, None)
self.set_magvit_padding_more_frame()
new_pixel_values = [first_frames]
for i in range(1, x.shape[2], self.mini_batch_encoder):
next_frames = self.single_forward(x[:, :, i: i + self.mini_batch_encoder, :, :], None, None)
new_pixel_values.append(next_frames)
new_pixel_values = torch.cat(new_pixel_values, dim=2)
elif self.cache_compression_vae:
_, _, f, _, _ = x.size()
if f % 2 != 0:
self.set_padding_one_frame()
first_frames = self.single_forward(x[:, :, 0:1, :, :], None, None)
self.set_padding_more_frame()
new_pixel_values = [first_frames]
start_index = 1
else:
self.set_padding_more_frame()
new_pixel_values = []
start_index = 0
for i in range(start_index, x.shape[2], self.mini_batch_encoder):
next_frames = self.single_forward(x[:, :, i: i + self.mini_batch_encoder, :, :], None, None)
new_pixel_values.append(next_frames)
new_pixel_values = torch.cat(new_pixel_values, dim=2)
elif self.slice_compression_vae:
_, _, f, _, _ = x.size()
if f % 2 != 0:
self.set_padding_one_frame()
first_frames = self.single_forward(x[:, :, 0:1, :, :], None, None)
self.set_padding_more_frame()
new_pixel_values = [first_frames]
start_index = 1
else:
self.set_padding_more_frame()
new_pixel_values = []
start_index = 0
for i in range(start_index, x.shape[2], self.mini_batch_encoder):
next_frames = self.single_forward(x[:, :, i: i + self.mini_batch_encoder, :, :], None, None)
new_pixel_values.append(next_frames)
new_pixel_values = torch.cat(new_pixel_values, dim=2)
elif self.slice_mag_vae:
_, _, f, _, _ = x.size()
new_pixel_values = []
for i in range(0, x.shape[2], self.mini_batch_encoder):
next_frames = self.single_forward(x[:, :, i: i + self.mini_batch_encoder, :, :], None, None)
new_pixel_values.append(next_frames)
new_pixel_values = torch.cat(new_pixel_values, dim=2)
else:
new_pixel_values = self.single_forward(x, None, None)
return new_pixel_values
class Decoder(nn.Module):
r"""
The `Decoder` layer of a variational autoencoder that decodes its latent representation into an output sample.
Args:
in_channels (`int`, *optional*, defaults to 8):
The number of input channels.
out_channels (`int`, *optional*, defaults to 3):
The number of output channels.
up_block_types (`Tuple[str, ...]`, *optional*, defaults to `("SpatialUpBlock3D",)`):
The types of up blocks to use.
block_out_channels (`Tuple[int, ...]`, *optional*, defaults to `(64,)`):
The number of output channels for each block.
use_gc_blocks (`Tuple[bool, ...]`, *optional*, defaults to `None`):
Whether to use global context blocks for each down block.
mid_block_type (`str`, *optional*, defaults to `"MidBlock3D"`):
The type of mid block to use.
layers_per_block (`int`, *optional*, defaults to 2):
The number of layers per block.
norm_num_groups (`int`, *optional*, defaults to 32):
The number of groups for normalization.
act_fn (`str`, *optional*, defaults to `"silu"`):
The activation function to use. See `~diffusers.models.activations.get_activation` for available options.
num_attention_heads (`int`, *optional*, defaults to 1):
The number of attention heads to use.
"""
def __init__(
self,
in_channels: int = 8,
out_channels: int = 3,
up_block_types = ("SpatialUpBlock3D",),
ch = 128,
ch_mult = [1,2,4,4,],
block_out_channels = [128, 256, 512, 512],
use_gc_blocks = None,
mid_block_type: str = "MidBlock3D",
mid_block_use_attention: bool = True,
mid_block_attention_type: str = "3d",
mid_block_num_attention_heads: int = 1,
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
num_attention_heads: int = 1,
slice_mag_vae: bool = False,
slice_compression_vae: bool = False,
cache_compression_vae: bool = False,
cache_mag_vae: bool = False,
spatial_group_norm: bool = False,
mini_batch_decoder: int = 3,
verbose = False,
):
super().__init__()
if block_out_channels is None:
block_out_channels = [ch * i for i in ch_mult]
assert len(up_block_types) == len(block_out_channels), (
"Number of up block types must match number of block output channels."
)
if use_gc_blocks is not None:
assert len(use_gc_blocks) == len(up_block_types), (
"Number of GC blocks must match number of up block types."
)
else:
use_gc_blocks = [False] * len(up_block_types)
self.conv_in = CausalConv3d(
in_channels,
block_out_channels[-1],
kernel_size=3,
)
self.mid_block = get_mid_block(
mid_block_type,
in_channels=block_out_channels[-1],
num_layers=layers_per_block,
act_fn=act_fn,
norm_num_groups=norm_num_groups,
norm_eps=1e-6,
add_attention=mid_block_use_attention,
attention_type=mid_block_attention_type,
num_attention_heads=mid_block_num_attention_heads,
)
self.up_blocks = nn.ModuleList([])
reversed_block_out_channels = list(reversed(block_out_channels))
output_channels = reversed_block_out_channels[0]
for i, up_block_type in enumerate(up_block_types):
input_channels = output_channels
output_channels = reversed_block_out_channels[i]
# is_first_block = i == 0
is_final_block = i == len(block_out_channels) - 1
up_block = get_up_block(
up_block_type,
in_channels=input_channels,
out_channels=output_channels,
num_layers=layers_per_block + 1,
act_fn=act_fn,
norm_num_groups=norm_num_groups,
norm_eps=1e-6,
num_attention_heads=num_attention_heads,
add_gc_block=use_gc_blocks[i],
add_upsample=not is_final_block,
)
self.up_blocks.append(up_block)
self.conv_norm_out = nn.GroupNorm(
num_channels=block_out_channels[0],
num_groups=norm_num_groups,
eps=1e-6,
)
self.conv_act = get_activation(act_fn)
self.conv_out = CausalConv3d(block_out_channels[0], out_channels, kernel_size=3)
self.slice_mag_vae = slice_mag_vae
self.slice_compression_vae = slice_compression_vae
self.cache_compression_vae = cache_compression_vae
self.cache_mag_vae = cache_mag_vae
self.mini_batch_decoder = mini_batch_decoder
self.spatial_group_norm = spatial_group_norm
self.verbose = verbose
def set_padding_one_frame(self):
def _set_padding_one_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 1
for sub_name, sub_mod in module.named_children():
_set_padding_one_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_padding_one_frame(name, module)
def set_padding_more_frame(self):
def _set_padding_more_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 2
for sub_name, sub_mod in module.named_children():
_set_padding_more_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_padding_more_frame(name, module)
def set_magvit_padding_one_frame(self):
def _set_magvit_padding_one_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 3
for sub_name, sub_mod in module.named_children():
_set_magvit_padding_one_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_magvit_padding_one_frame(name, module)
def set_magvit_padding_more_frame(self):
def _set_magvit_padding_more_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 4
for sub_name, sub_mod in module.named_children():
_set_magvit_padding_more_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_magvit_padding_more_frame(name, module)
def set_cache_slice_vae_padding_one_frame(self):
def _set_cache_slice_vae_padding_one_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 5
for sub_name, sub_mod in module.named_children():
_set_cache_slice_vae_padding_one_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_cache_slice_vae_padding_one_frame(name, module)
def set_cache_slice_vae_padding_more_frame(self):
def _set_cache_slice_vae_padding_more_frame(name, module):
if hasattr(module, 'padding_flag'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.padding_flag = 6
for sub_name, sub_mod in module.named_children():
_set_cache_slice_vae_padding_more_frame(sub_name, sub_mod)
for name, module in self.named_children():
_set_cache_slice_vae_padding_more_frame(name, module)
def set_3dgroupnorm_for_submodule(self):
def _set_3dgroupnorm_for_submodule(name, module):
if hasattr(module, 'set_3dgroupnorm'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.set_3dgroupnorm = True
for sub_name, sub_mod in module.named_children():
_set_3dgroupnorm_for_submodule(sub_name, sub_mod)
for name, module in self.named_children():
_set_3dgroupnorm_for_submodule(name, module)
def clear_cache(self):
def _clear_cache(name, module):
if hasattr(module, 'prev_features'):
if self.verbose:
print('Set pad mode for module[%s] type=%s' % (name, str(type(module))))
module.prev_features = None
for sub_name, sub_mod in module.named_children():
_clear_cache(sub_name, sub_mod)
for name, module in self.named_children():
_clear_cache(name, module)
def single_forward(self, x: torch.Tensor, previous_features: torch.Tensor, after_features: torch.Tensor) -> torch.Tensor:
# x: (B, C, T, H, W)
if self.training:
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
if previous_features is not None and after_features is None:
b, c, t, h, w = x.size()
x = torch.concat([previous_features, x], 2)
x = self.conv_in(x)
x = self.mid_block(x)
x = x[:, :, -t:]
elif previous_features is None and after_features is not None:
b, c, t, h, w = x.size()
x = torch.concat([x, after_features], 2)
x = self.conv_in(x)
x = self.mid_block(x)
x = x[:, :, :t]
elif previous_features is not None and after_features is not None:
_, _, t_1, _, _ = previous_features.size()
_, _, t_2, _, _ = x.size()
x = torch.concat([previous_features, x, after_features], 2)
x = self.conv_in(x)
x = self.mid_block(x)
x = x[:, :, t_1:(t_1 + t_2)]
else:
if self.training:
x = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.conv_in),
x,
**ckpt_kwargs,
)
x = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.mid_block),
x,
**ckpt_kwargs,
)
else:
x = self.conv_in(x)
x = self.mid_block(x)
for up_block in self.up_blocks:
if self.training:
x = torch.utils.checkpoint.checkpoint(
create_custom_forward(up_block),
x,
**ckpt_kwargs,
)
else:
x = up_block(x)
if self.spatial_group_norm:
batch_size = x.shape[0]
x = rearrange(x, "b c t h w -> (b t) c h w")
x = self.conv_norm_out(x)
x = rearrange(x, "(b t) c h w -> b c t h w", b=batch_size)
else:
x = self.conv_norm_out(x)
x = self.conv_act(x)
x = self.conv_out(x)
return x
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.spatial_group_norm:
self.set_3dgroupnorm_for_submodule()
if self.cache_mag_vae:
self.set_magvit_padding_one_frame()
first_frames = self.single_forward(x[:, :, 0:1, :, :], None, None)
self.set_magvit_padding_more_frame()
new_pixel_values = [first_frames]
for i in range(1, x.shape[2], self.mini_batch_decoder):
next_frames = self.single_forward(x[:, :, i: i + self.mini_batch_decoder, :, :], None, None)
new_pixel_values.append(next_frames)
new_pixel_values = torch.cat(new_pixel_values, dim=2)
elif self.cache_compression_vae:
_, _, f, _, _ = x.size()
if f == 1:
self.set_padding_one_frame()
first_frames = self.single_forward(x[:, :, :1, :, :], None, None)
new_pixel_values = [first_frames]
start_index = 1
else:
self.set_cache_slice_vae_padding_one_frame()
first_frames = self.single_forward(x[:, :, :self.mini_batch_decoder, :, :], None, None)
new_pixel_values = [first_frames]
start_index = self.mini_batch_decoder
for i in range(start_index, x.shape[2], self.mini_batch_decoder):
self.set_cache_slice_vae_padding_more_frame()
next_frames = self.single_forward(x[:, :, i: i + self.mini_batch_decoder, :, :], None, None)
new_pixel_values.append(next_frames)
new_pixel_values = torch.cat(new_pixel_values, dim=2)
elif self.slice_compression_vae:
_, _, f, _, _ = x.size()
if f % 2 != 0:
self.set_padding_one_frame()
first_frames = self.single_forward(x[:, :, 0:1, :, :], None, None)
self.set_padding_more_frame()
new_pixel_values = [first_frames]
start_index = 1
else:
self.set_padding_more_frame()
new_pixel_values = []
start_index = 0
previous_features = None
for i in range(start_index, x.shape[2], self.mini_batch_decoder):
after_features = x[:, :, i + self.mini_batch_decoder: i + 2 * self.mini_batch_decoder, :, :] if i + self.mini_batch_decoder < x.shape[2] else None
next_frames = self.single_forward(x[:, :, i: i + self.mini_batch_decoder, :, :], previous_features, after_features)
previous_features = x[:, :, i: i + self.mini_batch_decoder, :, :]
new_pixel_values.append(next_frames)
new_pixel_values = torch.cat(new_pixel_values, dim=2)
elif self.slice_mag_vae:
_, _, f, _, _ = x.size()
new_pixel_values = []
for i in range(0, x.shape[2], self.mini_batch_decoder):
next_frames = self.single_forward(x[:, :, i: i + self.mini_batch_decoder, :, :], None, None)
new_pixel_values.append(next_frames)
new_pixel_values = torch.cat(new_pixel_values, dim=2)
else:
new_pixel_values = self.single_forward(x, None, None)
return new_pixel_values
@@ -0,0 +1,701 @@
# pytorch_diffusion + derived encoder decoder
import math
import numpy as np
import torch
import torch.nn as nn
from einops import rearrange
from ...util import instantiate_from_config
def get_timestep_embedding(timesteps, embedding_dim):
"""
This matches the implementation in Denoising Diffusion Probabilistic Models:
From Fairseq.
Build sinusoidal embeddings.
This matches the implementation in tensor2tensor, but differs slightly
from the description in Section 3.5 of "Attention Is All You Need".
"""
assert len(timesteps.shape) == 1
half_dim = embedding_dim // 2
emb = math.log(10000) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, dtype=torch.float32) * -emb)
emb = emb.to(device=timesteps.device)
emb = timesteps.float()[:, None] * emb[None, :]
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1)
if embedding_dim % 2 == 1: # zero pad
emb = torch.nn.functional.pad(emb, (0,1,0,0))
return emb
def nonlinearity(x):
# swish
return x*torch.sigmoid(x)
def Normalize(in_channels, num_groups=32):
return torch.nn.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True)
class Upsample(nn.Module):
def __init__(self, in_channels, with_conv):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
self.conv = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=3,
stride=1,
padding=1)
def forward(self, x):
x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
if self.with_conv:
x = self.conv(x)
return x
class Downsample(nn.Module):
def __init__(self, in_channels, with_conv):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
# no asymmetric padding in torch conv, must do it ourselves
self.conv = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=3,
stride=2,
padding=0)
def forward(self, x):
if self.with_conv:
pad = (0,1,0,1)
x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
x = self.conv(x)
else:
x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2)
return x
class ResnetBlock(nn.Module):
def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False,
dropout, temb_channels=512):
super().__init__()
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.use_conv_shortcut = conv_shortcut
self.norm1 = Normalize(in_channels)
self.conv1 = torch.nn.Conv2d(in_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1)
if temb_channels > 0:
self.temb_proj = torch.nn.Linear(temb_channels,
out_channels)
self.norm2 = Normalize(out_channels)
self.dropout = torch.nn.Dropout(dropout)
self.conv2 = torch.nn.Conv2d(out_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1)
if self.in_channels != self.out_channels:
if self.use_conv_shortcut:
self.conv_shortcut = torch.nn.Conv2d(in_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1)
else:
self.nin_shortcut = torch.nn.Conv2d(in_channels,
out_channels,
kernel_size=1,
stride=1,
padding=0)
def forward(self, x, temb):
h = x
h = self.norm1(h)
h = nonlinearity(h)
h = self.conv1(h)
if temb is not None:
h = h + self.temb_proj(nonlinearity(temb))[:,:,None,None]
h = self.norm2(h)
h = nonlinearity(h)
h = self.dropout(h)
h = self.conv2(h)
if self.in_channels != self.out_channels:
if self.use_conv_shortcut:
x = self.conv_shortcut(x)
else:
x = self.nin_shortcut(x)
return x+h
class AttnBlock(nn.Module):
def __init__(self, in_channels):
super().__init__()
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.k = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.v = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.proj_out = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
def forward(self, x):
h_ = x
h_ = self.norm(h_)
q = self.q(h_)
k = self.k(h_)
v = self.v(h_)
# compute attention
b,c,h,w = q.shape
q = q.reshape(b,c,h*w)
q = q.permute(0,2,1) # b,hw,c
k = k.reshape(b,c,h*w) # b,c,hw
w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
w_ = w_ * (int(c)**(-0.5))
w_ = torch.nn.functional.softmax(w_, dim=2)
# attend to values
v = v.reshape(b,c,h*w)
w_ = w_.permute(0,2,1) # b,hw,hw (first hw of k, second of q)
h_ = torch.bmm(v,w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
h_ = h_.reshape(b,c,h,w)
h_ = self.proj_out(h_)
return x+h_
class LinearAttention(nn.Module):
def __init__(self, dim, heads=4, dim_head=32):
super().__init__()
self.heads = heads
hidden_dim = dim_head * heads
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
self.to_out = nn.Conv2d(hidden_dim, dim, 1)
def forward(self, x):
b, c, h, w = x.shape
qkv = self.to_qkv(x)
q, k, v = rearrange(qkv, 'b (qkv heads c) h w -> qkv b heads c (h w)', heads = self.heads, qkv=3)
k = k.softmax(dim=-1)
context = torch.einsum('bhdn,bhen->bhde', k, v)
out = torch.einsum('bhde,bhdn->bhen', context, q)
out = rearrange(out, 'b heads c (h w) -> b (heads c) h w', heads=self.heads, h=h, w=w)
return self.to_out(out)
class LinAttnBlock(LinearAttention):
"""to match AttnBlock usage"""
def __init__(self, in_channels):
super().__init__(dim=in_channels, heads=1, dim_head=in_channels)
def make_attn(in_channels, attn_type="vanilla"):
assert attn_type in ["vanilla", "linear", "none"], f'attn_type {attn_type} unknown'
print(f"making attention of type '{attn_type}' with {in_channels} in_channels")
if attn_type == "vanilla":
return AttnBlock(in_channels)
elif attn_type == "none":
return nn.Identity(in_channels)
else:
return LinAttnBlock(in_channels)
class Encoder(nn.Module):
def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
resolution, z_channels, double_z=True, use_linear_attn=False, attn_type="vanilla",
**ignore_kwargs):
super().__init__()
if use_linear_attn: attn_type = "linear"
self.ch = ch
self.temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.in_channels = in_channels
# downsampling
self.conv_in = torch.nn.Conv2d(in_channels,
self.ch,
kernel_size=3,
stride=1,
padding=1)
curr_res = resolution
in_ch_mult = (1,)+tuple(ch_mult)
self.in_ch_mult = in_ch_mult
self.down = nn.ModuleList()
for i_level in range(self.num_resolutions):
block = nn.ModuleList()
attn = nn.ModuleList()
block_in = ch*in_ch_mult[i_level]
block_out = ch*ch_mult[i_level]
for i_block in range(self.num_res_blocks):
block.append(ResnetBlock(in_channels=block_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout))
block_in = block_out
if curr_res in attn_resolutions:
attn.append(make_attn(block_in, attn_type=attn_type))
down = nn.Module()
down.block = block
down.attn = attn
if i_level != self.num_resolutions-1:
down.downsample = Downsample(block_in, resamp_with_conv)
curr_res = curr_res // 2
self.down.append(down)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock(in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
self.mid.attn_1 = make_attn(block_in, attn_type=attn_type)
self.mid.block_2 = ResnetBlock(in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
2*z_channels if double_z else z_channels,
kernel_size=3,
stride=1,
padding=1)
def forward(self, x):
# timestep embedding
temb = None
# downsampling
hs = [self.conv_in(x)]
for i_level in range(self.num_resolutions):
for i_block in range(self.num_res_blocks):
h = self.down[i_level].block[i_block](hs[-1], temb)
if len(self.down[i_level].attn) > 0:
h = self.down[i_level].attn[i_block](h)
hs.append(h)
if i_level != self.num_resolutions-1:
hs.append(self.down[i_level].downsample(hs[-1]))
# middle
h = hs[-1]
h = self.mid.block_1(h, temb)
h = self.mid.attn_1(h)
h = self.mid.block_2(h, temb)
# end
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
return h
class Decoder(nn.Module):
def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
resolution, z_channels, give_pre_end=False, tanh_out=False, use_linear_attn=False,
attn_type="vanilla", **ignorekwargs):
super().__init__()
if use_linear_attn: attn_type = "linear"
self.ch = ch
self.temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.in_channels = in_channels
self.give_pre_end = give_pre_end
self.tanh_out = tanh_out
# compute in_ch_mult, block_in and curr_res at lowest res
in_ch_mult = (1,)+tuple(ch_mult)
block_in = ch*ch_mult[self.num_resolutions-1]
curr_res = resolution // 2**(self.num_resolutions-1)
self.z_shape = (1,z_channels,curr_res,curr_res)
print("Working with z of shape {} = {} dimensions.".format(
self.z_shape, np.prod(self.z_shape)))
# z to block_in
self.conv_in = torch.nn.Conv2d(z_channels,
block_in,
kernel_size=3,
stride=1,
padding=1)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock(in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
self.mid.attn_1 = make_attn(block_in, attn_type=attn_type)
self.mid.block_2 = ResnetBlock(in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
# upsampling
self.up = nn.ModuleList()
for i_level in reversed(range(self.num_resolutions)):
block = nn.ModuleList()
attn = nn.ModuleList()
block_out = ch*ch_mult[i_level]
for i_block in range(self.num_res_blocks+1):
block.append(ResnetBlock(in_channels=block_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout))
block_in = block_out
if curr_res in attn_resolutions:
attn.append(make_attn(block_in, attn_type=attn_type))
up = nn.Module()
up.block = block
up.attn = attn
if i_level != 0:
up.upsample = Upsample(block_in, resamp_with_conv)
curr_res = curr_res * 2
self.up.insert(0, up) # prepend to get consistent order
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
out_ch,
kernel_size=3,
stride=1,
padding=1)
def forward(self, z):
#assert z.shape[1:] == self.z_shape[1:]
self.last_z_shape = z.shape
# timestep embedding
temb = None
# z to block_in
h = self.conv_in(z)
# middle
h = self.mid.block_1(h, temb)
h = self.mid.attn_1(h)
h = self.mid.block_2(h, temb)
# upsampling
for i_level in reversed(range(self.num_resolutions)):
for i_block in range(self.num_res_blocks+1):
h = self.up[i_level].block[i_block](h, temb)
if len(self.up[i_level].attn) > 0:
h = self.up[i_level].attn[i_block](h)
if i_level != 0:
h = self.up[i_level].upsample(h)
# end
if self.give_pre_end:
return h
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
if self.tanh_out:
h = torch.tanh(h)
return h
class SimpleDecoder(nn.Module):
def __init__(self, in_channels, out_channels, *args, **kwargs):
super().__init__()
self.model = nn.ModuleList([nn.Conv2d(in_channels, in_channels, 1),
ResnetBlock(in_channels=in_channels,
out_channels=2 * in_channels,
temb_channels=0, dropout=0.0),
ResnetBlock(in_channels=2 * in_channels,
out_channels=4 * in_channels,
temb_channels=0, dropout=0.0),
ResnetBlock(in_channels=4 * in_channels,
out_channels=2 * in_channels,
temb_channels=0, dropout=0.0),
nn.Conv2d(2*in_channels, in_channels, 1),
Upsample(in_channels, with_conv=True)])
# end
self.norm_out = Normalize(in_channels)
self.conv_out = torch.nn.Conv2d(in_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1)
def forward(self, x):
for i, layer in enumerate(self.model):
if i in [1,2,3]:
x = layer(x, None)
else:
x = layer(x)
h = self.norm_out(x)
h = nonlinearity(h)
x = self.conv_out(h)
return x
class UpsampleDecoder(nn.Module):
def __init__(self, in_channels, out_channels, ch, num_res_blocks, resolution,
ch_mult=(2,2), dropout=0.0):
super().__init__()
# upsampling
self.temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
block_in = in_channels
curr_res = resolution // 2 ** (self.num_resolutions - 1)
self.res_blocks = nn.ModuleList()
self.upsample_blocks = nn.ModuleList()
for i_level in range(self.num_resolutions):
res_block = []
block_out = ch * ch_mult[i_level]
for i_block in range(self.num_res_blocks + 1):
res_block.append(ResnetBlock(in_channels=block_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout))
block_in = block_out
self.res_blocks.append(nn.ModuleList(res_block))
if i_level != self.num_resolutions - 1:
self.upsample_blocks.append(Upsample(block_in, True))
curr_res = curr_res * 2
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
out_channels,
kernel_size=3,
stride=1,
padding=1)
def forward(self, x):
# upsampling
h = x
for k, i_level in enumerate(range(self.num_resolutions)):
for i_block in range(self.num_res_blocks + 1):
h = self.res_blocks[i_level][i_block](h, None)
if i_level != self.num_resolutions - 1:
h = self.upsample_blocks[k](h)
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
return h
class LatentRescaler(nn.Module):
def __init__(self, factor, in_channels, mid_channels, out_channels, depth=2):
super().__init__()
# residual block, interpolate, residual block
self.factor = factor
self.conv_in = nn.Conv2d(in_channels,
mid_channels,
kernel_size=3,
stride=1,
padding=1)
self.res_block1 = nn.ModuleList([ResnetBlock(in_channels=mid_channels,
out_channels=mid_channels,
temb_channels=0,
dropout=0.0) for _ in range(depth)])
self.attn = AttnBlock(mid_channels)
self.res_block2 = nn.ModuleList([ResnetBlock(in_channels=mid_channels,
out_channels=mid_channels,
temb_channels=0,
dropout=0.0) for _ in range(depth)])
self.conv_out = nn.Conv2d(mid_channels,
out_channels,
kernel_size=1,
)
def forward(self, x):
x = self.conv_in(x)
for block in self.res_block1:
x = block(x, None)
x = torch.nn.functional.interpolate(x, size=(int(round(x.shape[2]*self.factor)), int(round(x.shape[3]*self.factor))))
x = self.attn(x)
for block in self.res_block2:
x = block(x, None)
x = self.conv_out(x)
return x
class MergedRescaleEncoder(nn.Module):
def __init__(self, in_channels, ch, resolution, out_ch, num_res_blocks,
attn_resolutions, dropout=0.0, resamp_with_conv=True,
ch_mult=(1,2,4,8), rescale_factor=1.0, rescale_module_depth=1):
super().__init__()
intermediate_chn = ch * ch_mult[-1]
self.encoder = Encoder(in_channels=in_channels, num_res_blocks=num_res_blocks, ch=ch, ch_mult=ch_mult,
z_channels=intermediate_chn, double_z=False, resolution=resolution,
attn_resolutions=attn_resolutions, dropout=dropout, resamp_with_conv=resamp_with_conv,
out_ch=None)
self.rescaler = LatentRescaler(factor=rescale_factor, in_channels=intermediate_chn,
mid_channels=intermediate_chn, out_channels=out_ch, depth=rescale_module_depth)
def forward(self, x):
x = self.encoder(x)
x = self.rescaler(x)
return x
class MergedRescaleDecoder(nn.Module):
def __init__(self, z_channels, out_ch, resolution, num_res_blocks, attn_resolutions, ch, ch_mult=(1,2,4,8),
dropout=0.0, resamp_with_conv=True, rescale_factor=1.0, rescale_module_depth=1):
super().__init__()
tmp_chn = z_channels*ch_mult[-1]
self.decoder = Decoder(out_ch=out_ch, z_channels=tmp_chn, attn_resolutions=attn_resolutions, dropout=dropout,
resamp_with_conv=resamp_with_conv, in_channels=None, num_res_blocks=num_res_blocks,
ch_mult=ch_mult, resolution=resolution, ch=ch)
self.rescaler = LatentRescaler(factor=rescale_factor, in_channels=z_channels, mid_channels=tmp_chn,
out_channels=tmp_chn, depth=rescale_module_depth)
def forward(self, x):
x = self.rescaler(x)
x = self.decoder(x)
return x
class Upsampler(nn.Module):
def __init__(self, in_size, out_size, in_channels, out_channels, ch_mult=2):
super().__init__()
assert out_size >= in_size
num_blocks = int(np.log2(out_size//in_size))+1
factor_up = 1.+ (out_size % in_size)
print(f"Building {self.__class__.__name__} with in_size: {in_size} --> out_size {out_size} and factor {factor_up}")
self.rescaler = LatentRescaler(factor=factor_up, in_channels=in_channels, mid_channels=2*in_channels,
out_channels=in_channels)
self.decoder = Decoder(out_ch=out_channels, resolution=out_size, z_channels=in_channels, num_res_blocks=2,
attn_resolutions=[], in_channels=None, ch=in_channels,
ch_mult=[ch_mult for _ in range(num_blocks)])
def forward(self, x):
x = self.rescaler(x)
x = self.decoder(x)
return x
class Resize(nn.Module):
def __init__(self, in_channels=None, learned=False, mode="bilinear"):
super().__init__()
self.with_conv = learned
self.mode = mode
if self.with_conv:
print(f"Note: {self.__class__.__name} uses learned downsampling and will ignore the fixed {mode} mode")
raise NotImplementedError()
assert in_channels is not None
# no asymmetric padding in torch conv, must do it ourselves
self.conv = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=4,
stride=2,
padding=1)
def forward(self, x, scale_factor=1.0):
if scale_factor==1.0:
return x
else:
x = torch.nn.functional.interpolate(x, mode=self.mode, align_corners=False, scale_factor=scale_factor)
return x
class FirstStagePostProcessor(nn.Module):
def __init__(self, ch_mult:list, in_channels,
pretrained_model:nn.Module=None,
reshape=False,
n_channels=None,
dropout=0.,
pretrained_config=None):
super().__init__()
if pretrained_config is None:
assert pretrained_model is not None, 'Either "pretrained_model" or "pretrained_config" must not be None'
self.pretrained_model = pretrained_model
else:
assert pretrained_config is not None, 'Either "pretrained_model" or "pretrained_config" must not be None'
self.instantiate_pretrained(pretrained_config)
self.do_reshape = reshape
if n_channels is None:
n_channels = self.pretrained_model.encoder.ch
self.proj_norm = Normalize(in_channels,num_groups=in_channels//2)
self.proj = nn.Conv2d(in_channels,n_channels,kernel_size=3,
stride=1,padding=1)
blocks = []
downs = []
ch_in = n_channels
for m in ch_mult:
blocks.append(ResnetBlock(in_channels=ch_in,out_channels=m*n_channels,dropout=dropout))
ch_in = m * n_channels
downs.append(Downsample(ch_in, with_conv=False))
self.model = nn.ModuleList(blocks)
self.downsampler = nn.ModuleList(downs)
def instantiate_pretrained(self, config):
model = instantiate_from_config(config)
self.pretrained_model = model.eval()
# self.pretrained_model.train = False
for param in self.pretrained_model.parameters():
param.requires_grad = False
@torch.no_grad()
def encode_with_pretrained(self,x):
c = self.pretrained_model.encode(x)
if isinstance(c, DiagonalGaussianDistribution):
c = c.mode()
return c
def forward(self,x):
z_fs = self.encode_with_pretrained(x)
z = self.proj_norm(z_fs)
z = self.proj(z)
z = nonlinearity(z)
for submodel, downmodel in zip(self.model,self.downsampler):
z = submodel(z,temb=None)
z = downmodel(z)
if self.do_reshape:
z = rearrange(z,'b c h w -> b (h w) c')
return z
@@ -0,0 +1,268 @@
# adopted from
# https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py
# and
# https://github.com/lucidrains/denoising-diffusion-pytorch/blob/7706bdfc6f527f58d33f84b7b522e61e6e3164b3/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py
# and
# https://github.com/openai/guided-diffusion/blob/0ba878e517b276c45d1195eb29f6f5f72659a05b/guided_diffusion/nn.py
#
# thanks!
import math
import os
import numpy as np
import torch
import torch.nn as nn
from einops import repeat
from ...util import instantiate_from_config
def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3):
if schedule == "linear":
betas = (
torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=torch.float64) ** 2
)
elif schedule == "cosine":
timesteps = (
torch.arange(n_timestep + 1, dtype=torch.float64) / n_timestep + cosine_s
)
alphas = timesteps / (1 + cosine_s) * np.pi / 2
alphas = torch.cos(alphas).pow(2)
alphas = alphas / alphas[0]
betas = 1 - alphas[1:] / alphas[:-1]
betas = np.clip(betas, a_min=0, a_max=0.999)
elif schedule == "sqrt_linear":
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64)
elif schedule == "sqrt":
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64) ** 0.5
else:
raise ValueError(f"schedule '{schedule}' unknown.")
return betas.numpy()
def make_ddim_timesteps(ddim_discr_method, num_ddim_timesteps, num_ddpm_timesteps, verbose=True):
if ddim_discr_method == 'uniform':
c = num_ddpm_timesteps // num_ddim_timesteps
ddim_timesteps = np.asarray(list(range(0, num_ddpm_timesteps, c)))
elif ddim_discr_method == 'quad':
ddim_timesteps = ((np.linspace(0, np.sqrt(num_ddpm_timesteps * .8), num_ddim_timesteps)) ** 2).astype(int)
else:
raise NotImplementedError(f'There is no ddim discretization method called "{ddim_discr_method}"')
# assert ddim_timesteps.shape[0] == num_ddim_timesteps
# add one to get the final alpha values right (the ones from first scale to data during sampling)
steps_out = ddim_timesteps + 1
if verbose:
print(f'Selected timesteps for ddim sampler: {steps_out}')
return steps_out
def make_ddim_sampling_parameters(alphacums, ddim_timesteps, eta, verbose=True):
# select alphas for computing the variance schedule
alphas = alphacums[ddim_timesteps]
alphas_prev = np.asarray([alphacums[0]] + alphacums[ddim_timesteps[:-1]].tolist())
# according the the formula provided in https://arxiv.org/abs/2010.02502
sigmas = eta * np.sqrt((1 - alphas_prev) / (1 - alphas) * (1 - alphas / alphas_prev))
if verbose:
print(f'Selected alphas for ddim sampler: a_t: {alphas}; a_(t-1): {alphas_prev}')
print(f'For the chosen value of eta, which is {eta}, '
f'this results in the following sigma_t schedule for ddim sampler {sigmas}')
return sigmas, alphas, alphas_prev
def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
"""
Create a beta schedule that discretizes the given alpha_t_bar function,
which defines the cumulative product of (1-beta) over time from t = [0,1].
:param num_diffusion_timesteps: the number of betas to produce.
:param alpha_bar: a lambda that takes an argument t from 0 to 1 and
produces the cumulative product of (1-beta) up to that
part of the diffusion process.
:param max_beta: the maximum beta to use; use values lower than 1 to
prevent singularities.
"""
betas = []
for i in range(num_diffusion_timesteps):
t1 = i / num_diffusion_timesteps
t2 = (i + 1) / num_diffusion_timesteps
betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta))
return np.array(betas)
def extract_into_tensor(a, t, x_shape):
b, *_ = t.shape
out = a.gather(-1, t)
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
def checkpoint(func, inputs, params, flag):
"""
Evaluate a function without caching intermediate activations, allowing for
reduced memory at the expense of extra compute in the backward pass.
:param func: the function to evaluate.
:param inputs: the argument sequence to pass to `func`.
:param params: a sequence of parameters `func` depends on but does not
explicitly take as arguments.
:param flag: if False, disable gradient checkpointing.
"""
if flag:
args = tuple(inputs) + tuple(params)
return CheckpointFunction.apply(func, len(inputs), *args)
else:
return func(*inputs)
class CheckpointFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, run_function, length, *args):
ctx.run_function = run_function
ctx.input_tensors = list(args[:length])
ctx.input_params = list(args[length:])
with torch.no_grad():
output_tensors = ctx.run_function(*ctx.input_tensors)
return output_tensors
@staticmethod
def backward(ctx, *output_grads):
ctx.input_tensors = [x.detach().requires_grad_(True) for x in ctx.input_tensors]
with torch.enable_grad():
# Fixes a bug where the first op in run_function modifies the
# Tensor storage in place, which is not allowed for detach()'d
# Tensors.
shallow_copies = [x.view_as(x) for x in ctx.input_tensors]
output_tensors = ctx.run_function(*shallow_copies)
input_grads = torch.autograd.grad(
output_tensors,
ctx.input_tensors + ctx.input_params,
output_grads,
allow_unused=True,
)
del ctx.input_tensors
del ctx.input_params
del output_tensors
return (None, None) + input_grads
def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False):
"""
Create sinusoidal timestep embeddings.
:param timesteps: a 1-D Tensor of N indices, one per batch element.
These may be fractional.
:param dim: the dimension of the output.
:param max_period: controls the minimum frequency of the embeddings.
:return: an [N x dim] Tensor of positional embeddings.
"""
if not repeat_only:
half = dim // 2
freqs = torch.exp(
-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
).to(device=timesteps.device)
args = timesteps[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
else:
embedding = repeat(timesteps, 'b -> b d', d=dim)
return embedding
def zero_module(module):
"""
Zero out the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().zero_()
return module
def scale_module(module, scale):
"""
Scale the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().mul_(scale)
return module
def mean_flat(tensor):
"""
Take the mean over all non-batch dimensions.
"""
return tensor.mean(dim=list(range(1, len(tensor.shape))))
def normalization(channels):
"""
Make a standard normalization layer.
:param channels: number of input channels.
:return: an nn.Module for normalization.
"""
return GroupNorm32(32, channels)
# PyTorch 1.7 has SiLU, but we support PyTorch 1.5.
class SiLU(nn.Module):
def forward(self, x):
return x * torch.sigmoid(x)
class GroupNorm32(nn.GroupNorm):
def forward(self, x):
return super().forward(x.float()).type(x.dtype)
def conv_nd(dims, *args, **kwargs):
"""
Create a 1D, 2D, or 3D convolution module.
"""
if dims == 1:
return nn.Conv1d(*args, **kwargs)
elif dims == 2:
return nn.Conv2d(*args, **kwargs)
elif dims == 3:
return nn.Conv3d(*args, **kwargs)
raise ValueError(f"unsupported dimensions: {dims}")
def linear(*args, **kwargs):
"""
Create a linear module.
"""
return nn.Linear(*args, **kwargs)
def avg_pool_nd(dims, *args, **kwargs):
"""
Create a 1D, 2D, or 3D average pooling module.
"""
if dims == 1:
return nn.AvgPool1d(*args, **kwargs)
elif dims == 2:
return nn.AvgPool2d(*args, **kwargs)
elif dims == 3:
return nn.AvgPool3d(*args, **kwargs)
raise ValueError(f"unsupported dimensions: {dims}")
class HybridConditioner(nn.Module):
def __init__(self, c_concat_config, c_crossattn_config):
super().__init__()
self.concat_conditioner = instantiate_from_config(c_concat_config)
self.crossattn_conditioner = instantiate_from_config(c_crossattn_config)
def forward(self, c_concat, c_crossattn):
c_concat = self.concat_conditioner(c_concat)
c_crossattn = self.crossattn_conditioner(c_crossattn)
return {'c_concat': [c_concat], 'c_crossattn': [c_crossattn]}
def noise_like(shape, device, repeat=False):
repeat_noise = lambda: torch.randn((1, *shape[1:]), device=device).repeat(shape[0], *((1,) * (len(shape) - 1)))
noise = lambda: torch.randn(shape, device=device)
return repeat_noise() if repeat else noise()
@@ -0,0 +1,92 @@
import numpy as np
import torch
class AbstractDistribution:
def sample(self):
raise NotImplementedError()
def mode(self):
raise NotImplementedError()
class DiracDistribution(AbstractDistribution):
def __init__(self, value):
self.value = value
def sample(self):
return self.value
def mode(self):
return self.value
class DiagonalGaussianDistribution(object):
def __init__(self, parameters, deterministic=False):
self.parameters = parameters
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
self.deterministic = deterministic
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
if self.deterministic:
self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)
def sample(self):
x = self.mean + self.std * torch.randn(self.mean.shape).to(device=self.parameters.device)
return x
def kl(self, other=None):
if self.deterministic:
return torch.Tensor([0.])
else:
if other is None:
return 0.5 * torch.sum(torch.pow(self.mean, 2)
+ self.var - 1.0 - self.logvar,
dim=[1, 2, 3])
else:
return 0.5 * torch.sum(
torch.pow(self.mean - other.mean, 2) / other.var
+ self.var / other.var - 1.0 - self.logvar + other.logvar,
dim=[1, 2, 3])
def nll(self, sample, dims=[1,2,3]):
if self.deterministic:
return torch.Tensor([0.])
logtwopi = np.log(2.0 * np.pi)
return 0.5 * torch.sum(
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
dim=dims)
def mode(self):
return self.mean
def normal_kl(mean1, logvar1, mean2, logvar2):
"""
source: https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/losses.py#L12
Compute the KL divergence between two gaussians.
Shapes are automatically broadcasted, so batches can be compared to
scalars, among other use cases.
"""
tensor = None
for obj in (mean1, logvar1, mean2, logvar2):
if isinstance(obj, torch.Tensor):
tensor = obj
break
assert tensor is not None, "at least one argument must be a Tensor"
# Force variances to be Tensors. Broadcasting helps convert scalars to
# Tensors, but it does not work for torch.exp().
logvar1, logvar2 = [
x if isinstance(x, torch.Tensor) else torch.tensor(x).to(tensor)
for x in (logvar1, logvar2)
]
return 0.5 * (
-1.0
+ logvar2
- logvar1
+ torch.exp(logvar1 - logvar2)
+ ((mean1 - mean2) ** 2) * torch.exp(-logvar2)
)
+116
View File
@@ -0,0 +1,116 @@
#-*- encoding:utf-8 -*-
import torch
from pytorch_lightning.callbacks import Callback
from torch import nn
class LitEma(nn.Module):
def __init__(self, model, decay=0.9999, use_num_upates=True):
super().__init__()
if decay < 0.0 or decay > 1.0:
raise ValueError('Decay must be between 0 and 1')
self.m_name2s_name = {}
self.register_buffer('decay', torch.tensor(decay, dtype=torch.float32))
self.register_buffer('num_updates', torch.tensor(0,dtype=torch.int) if use_num_upates
else torch.tensor(-1,dtype=torch.int))
for name, p in model.named_parameters():
if p.requires_grad:
#remove as '.'-character is not allowed in buffers
s_name = name.replace('.','')
self.m_name2s_name.update({name:s_name})
self.register_buffer(s_name,p.clone().detach().data)
self.collected_params = []
def forward(self,model):
decay = self.decay
if self.num_updates >= 0:
self.num_updates += 1
decay = min(self.decay,(1 + self.num_updates) / (10 + self.num_updates))
one_minus_decay = 1.0 - decay
with torch.no_grad():
m_param = dict(model.named_parameters())
shadow_params = dict(self.named_buffers())
for key in m_param:
if m_param[key].requires_grad:
sname = self.m_name2s_name[key]
shadow_params[sname] = shadow_params[sname].type_as(m_param[key])
shadow_params[sname].sub_(one_minus_decay * (shadow_params[sname] - m_param[key]))
else:
assert not key in self.m_name2s_name
def copy_to(self, model):
m_param = dict(model.named_parameters())
shadow_params = dict(self.named_buffers())
for key in m_param:
if m_param[key].requires_grad:
m_param[key].data.copy_(shadow_params[self.m_name2s_name[key]].data)
else:
assert not key in self.m_name2s_name
def store(self, parameters):
"""
Save the current parameters for restoring later.
Args:
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
temporarily stored.
"""
self.collected_params = [param.clone() for param in parameters]
def restore(self, parameters):
"""
Restore the parameters stored with the `store` method.
Useful to validate the model with EMA parameters without affecting the
original optimization process. Store the parameters before the
`copy_to` method. After validation (or model saving), use this to
restore the former parameters.
Args:
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
updated with the stored parameters.
"""
for c_param, param in zip(self.collected_params, parameters):
param.data.copy_(c_param.data)
class EMACallback(Callback):
def __init__(self, decay=0.9999):
self.decay = decay
self.shadow_params = {}
def on_train_start(self, trainer, pl_module):
# initialize shadow parameters for original models
total_ema_cnt = 0
for name, param in pl_module.named_parameters():
if name not in self.shadow_params:
self.shadow_params[name] = param.data.clone()
else: # already in dict, maybe load from checkpoint
pass
print('will calc ema for param: %s' % name)
total_ema_cnt += 1
print('total_ema_cnt=%d' % total_ema_cnt)
def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx):
# Update the shadow params at the end of each epoch
for name, param in pl_module.named_parameters():
assert name in self.shadow_params
new_average = (1.0 - self.decay) * param.data + self.decay * self.shadow_params[name]
self.shadow_params[name] = new_average.clone()
def on_save_checkpoint(self, trainer, pl_module, checkpoint):
# Save EMA parameters in the checkpoint
checkpoint['ema_params'] = self.shadow_params
def on_load_checkpoint(self, trainer, pl_module, checkpoint):
# Restore EMA parameters from the checkpoint
if 'ema_params' in checkpoint:
self.shadow_params = checkpoint.get('ema_params', {})
for k in self.shadow_params:
self.shadow_params[k] = self.shadow_params[k].cuda()
print('load shadow params from checkpoint, cnt=%d' % len(self.shadow_params))
else:
print('ema_params is not in checkpoint')
@@ -0,0 +1,3 @@
from .bsrgan import degradation_bsrgan_variant as degradation_fn_bsr
from .bsrgan_light import \
degradation_bsrgan_variant as degradation_fn_bsr_light
@@ -0,0 +1,730 @@
# -*- coding: utf-8 -*-
"""
# --------------------------------------------
# Super-Resolution
# --------------------------------------------
#
# Kai Zhang (cskaizhang@gmail.com)
# https://github.com/cszn
# From 2019/03--2021/08
# --------------------------------------------
"""
import random
from functools import partial
import albumentations
import cv2
import numpy as np
import scipy
import scipy.stats as ss
import torch
from scipy import ndimage
from scipy.interpolate import interp2d
from scipy.linalg import orth
from . import utils_image as util
def modcrop_np(img, sf):
'''
Args:
img: numpy image, WxH or WxHxC
sf: scale factor
Return:
cropped image
'''
w, h = img.shape[:2]
im = np.copy(img)
return im[:w - w % sf, :h - h % sf, ...]
"""
# --------------------------------------------
# anisotropic Gaussian kernels
# --------------------------------------------
"""
def analytic_kernel(k):
"""Calculate the X4 kernel from the X2 kernel (for proof see appendix in paper)"""
k_size = k.shape[0]
# Calculate the big kernels size
big_k = np.zeros((3 * k_size - 2, 3 * k_size - 2))
# Loop over the small kernel to fill the big one
for r in range(k_size):
for c in range(k_size):
big_k[2 * r:2 * r + k_size, 2 * c:2 * c + k_size] += k[r, c] * k
# Crop the edges of the big kernel to ignore very small values and increase run time of SR
crop = k_size // 2
cropped_big_k = big_k[crop:-crop, crop:-crop]
# Normalize to 1
return cropped_big_k / cropped_big_k.sum()
def anisotropic_Gaussian(ksize=15, theta=np.pi, l1=6, l2=6):
""" generate an anisotropic Gaussian kernel
Args:
ksize : e.g., 15, kernel size
theta : [0, pi], rotation angle range
l1 : [0.1,50], scaling of eigenvalues
l2 : [0.1,l1], scaling of eigenvalues
If l1 = l2, will get an isotropic Gaussian kernel.
Returns:
k : kernel
"""
v = np.dot(np.array([[np.cos(theta), -np.sin(theta)], [np.sin(theta), np.cos(theta)]]), np.array([1., 0.]))
V = np.array([[v[0], v[1]], [v[1], -v[0]]])
D = np.array([[l1, 0], [0, l2]])
Sigma = np.dot(np.dot(V, D), np.linalg.inv(V))
k = gm_blur_kernel(mean=[0, 0], cov=Sigma, size=ksize)
return k
def gm_blur_kernel(mean, cov, size=15):
center = size / 2.0 + 0.5
k = np.zeros([size, size])
for y in range(size):
for x in range(size):
cy = y - center + 1
cx = x - center + 1
k[y, x] = ss.multivariate_normal.pdf([cx, cy], mean=mean, cov=cov)
k = k / np.sum(k)
return k
def shift_pixel(x, sf, upper_left=True):
"""shift pixel for super-resolution with different scale factors
Args:
x: WxHxC or WxH
sf: scale factor
upper_left: shift direction
"""
h, w = x.shape[:2]
shift = (sf - 1) * 0.5
xv, yv = np.arange(0, w, 1.0), np.arange(0, h, 1.0)
if upper_left:
x1 = xv + shift
y1 = yv + shift
else:
x1 = xv - shift
y1 = yv - shift
x1 = np.clip(x1, 0, w - 1)
y1 = np.clip(y1, 0, h - 1)
if x.ndim == 2:
x = interp2d(xv, yv, x)(x1, y1)
if x.ndim == 3:
for i in range(x.shape[-1]):
x[:, :, i] = interp2d(xv, yv, x[:, :, i])(x1, y1)
return x
def blur(x, k):
'''
x: image, NxcxHxW
k: kernel, Nx1xhxw
'''
n, c = x.shape[:2]
p1, p2 = (k.shape[-2] - 1) // 2, (k.shape[-1] - 1) // 2
x = torch.nn.functional.pad(x, pad=(p1, p2, p1, p2), mode='replicate')
k = k.repeat(1, c, 1, 1)
k = k.view(-1, 1, k.shape[2], k.shape[3])
x = x.view(1, -1, x.shape[2], x.shape[3])
x = torch.nn.functional.conv2d(x, k, bias=None, stride=1, padding=0, groups=n * c)
x = x.view(n, c, x.shape[2], x.shape[3])
return x
def gen_kernel(k_size=np.array([15, 15]), scale_factor=np.array([4, 4]), min_var=0.6, max_var=10., noise_level=0):
""""
# modified version of https://github.com/assafshocher/BlindSR_dataset_generator
# Kai Zhang
# min_var = 0.175 * sf # variance of the gaussian kernel will be sampled between min_var and max_var
# max_var = 2.5 * sf
"""
# Set random eigen-vals (lambdas) and angle (theta) for COV matrix
lambda_1 = min_var + np.random.rand() * (max_var - min_var)
lambda_2 = min_var + np.random.rand() * (max_var - min_var)
theta = np.random.rand() * np.pi # random theta
noise = -noise_level + np.random.rand(*k_size) * noise_level * 2
# Set COV matrix using Lambdas and Theta
LAMBDA = np.diag([lambda_1, lambda_2])
Q = np.array([[np.cos(theta), -np.sin(theta)],
[np.sin(theta), np.cos(theta)]])
SIGMA = Q @ LAMBDA @ Q.T
INV_SIGMA = np.linalg.inv(SIGMA)[None, None, :, :]
# Set expectation position (shifting kernel for aligned image)
MU = k_size // 2 - 0.5 * (scale_factor - 1) # - 0.5 * (scale_factor - k_size % 2)
MU = MU[None, None, :, None]
# Create meshgrid for Gaussian
[X, Y] = np.meshgrid(range(k_size[0]), range(k_size[1]))
Z = np.stack([X, Y], 2)[:, :, :, None]
# Calcualte Gaussian for every pixel of the kernel
ZZ = Z - MU
ZZ_t = ZZ.transpose(0, 1, 3, 2)
raw_kernel = np.exp(-0.5 * np.squeeze(ZZ_t @ INV_SIGMA @ ZZ)) * (1 + noise)
# shift the kernel so it will be centered
# raw_kernel_centered = kernel_shift(raw_kernel, scale_factor)
# Normalize the kernel and return
# kernel = raw_kernel_centered / np.sum(raw_kernel_centered)
kernel = raw_kernel / np.sum(raw_kernel)
return kernel
def fspecial_gaussian(hsize, sigma):
hsize = [hsize, hsize]
siz = [(hsize[0] - 1.0) / 2.0, (hsize[1] - 1.0) / 2.0]
std = sigma
[x, y] = np.meshgrid(np.arange(-siz[1], siz[1] + 1), np.arange(-siz[0], siz[0] + 1))
arg = -(x * x + y * y) / (2 * std * std)
h = np.exp(arg)
h[h < scipy.finfo(float).eps * h.max()] = 0
sumh = h.sum()
if sumh != 0:
h = h / sumh
return h
def fspecial_laplacian(alpha):
alpha = max([0, min([alpha, 1])])
h1 = alpha / (alpha + 1)
h2 = (1 - alpha) / (alpha + 1)
h = [[h1, h2, h1], [h2, -4 / (alpha + 1), h2], [h1, h2, h1]]
h = np.array(h)
return h
def fspecial(filter_type, *args, **kwargs):
'''
python code from:
https://github.com/ronaldosena/imagens-medicas-2/blob/40171a6c259edec7827a6693a93955de2bd39e76/Aulas/aula_2_-_uniform_filter/matlab_fspecial.py
'''
if filter_type == 'gaussian':
return fspecial_gaussian(*args, **kwargs)
if filter_type == 'laplacian':
return fspecial_laplacian(*args, **kwargs)
"""
# --------------------------------------------
# degradation models
# --------------------------------------------
"""
def bicubic_degradation(x, sf=3):
'''
Args:
x: HxWxC image, [0, 1]
sf: down-scale factor
Return:
bicubicly downsampled LR image
'''
x = util.imresize_np(x, scale=1 / sf)
return x
def srmd_degradation(x, k, sf=3):
''' blur + bicubic downsampling
Args:
x: HxWxC image, [0, 1]
k: hxw, double
sf: down-scale factor
Return:
downsampled LR image
Reference:
@inproceedings{zhang2018learning,
title={Learning a single convolutional super-resolution network for multiple degradations},
author={Zhang, Kai and Zuo, Wangmeng and Zhang, Lei},
booktitle={IEEE Conference on Computer Vision and Pattern Recognition},
pages={3262--3271},
year={2018}
}
'''
x = ndimage.filters.convolve(x, np.expand_dims(k, axis=2), mode='wrap') # 'nearest' | 'mirror'
x = bicubic_degradation(x, sf=sf)
return x
def dpsr_degradation(x, k, sf=3):
''' bicubic downsampling + blur
Args:
x: HxWxC image, [0, 1]
k: hxw, double
sf: down-scale factor
Return:
downsampled LR image
Reference:
@inproceedings{zhang2019deep,
title={Deep Plug-and-Play Super-Resolution for Arbitrary Blur Kernels},
author={Zhang, Kai and Zuo, Wangmeng and Zhang, Lei},
booktitle={IEEE Conference on Computer Vision and Pattern Recognition},
pages={1671--1681},
year={2019}
}
'''
x = bicubic_degradation(x, sf=sf)
x = ndimage.filters.convolve(x, np.expand_dims(k, axis=2), mode='wrap')
return x
def classical_degradation(x, k, sf=3):
''' blur + downsampling
Args:
x: HxWxC image, [0, 1]/[0, 255]
k: hxw, double
sf: down-scale factor
Return:
downsampled LR image
'''
x = ndimage.filters.convolve(x, np.expand_dims(k, axis=2), mode='wrap')
# x = filters.correlate(x, np.expand_dims(np.flip(k), axis=2))
st = 0
return x[st::sf, st::sf, ...]
def add_sharpening(img, weight=0.5, radius=50, threshold=10):
"""USM sharpening. borrowed from real-ESRGAN
Input image: I; Blurry image: B.
1. K = I + weight * (I - B)
2. Mask = 1 if abs(I - B) > threshold, else: 0
3. Blur mask:
4. Out = Mask * K + (1 - Mask) * I
Args:
img (Numpy array): Input image, HWC, BGR; float32, [0, 1].
weight (float): Sharp weight. Default: 1.
radius (float): Kernel size of Gaussian blur. Default: 50.
threshold (int):
"""
if radius % 2 == 0:
radius += 1
blur = cv2.GaussianBlur(img, (radius, radius), 0)
residual = img - blur
mask = np.abs(residual) * 255 > threshold
mask = mask.astype('float32')
soft_mask = cv2.GaussianBlur(mask, (radius, radius), 0)
K = img + weight * residual
K = np.clip(K, 0, 1)
return soft_mask * K + (1 - soft_mask) * img
def add_blur(img, sf=4):
wd2 = 4.0 + sf
wd = 2.0 + 0.2 * sf
if random.random() < 0.5:
l1 = wd2 * random.random()
l2 = wd2 * random.random()
k = anisotropic_Gaussian(ksize=2 * random.randint(2, 11) + 3, theta=random.random() * np.pi, l1=l1, l2=l2)
else:
k = fspecial('gaussian', 2 * random.randint(2, 11) + 3, wd * random.random())
img = ndimage.filters.convolve(img, np.expand_dims(k, axis=2), mode='mirror')
return img
def add_resize(img, sf=4):
rnum = np.random.rand()
if rnum > 0.8: # up
sf1 = random.uniform(1, 2)
elif rnum < 0.7: # down
sf1 = random.uniform(0.5 / sf, 1)
else:
sf1 = 1.0
img = cv2.resize(img, (int(sf1 * img.shape[1]), int(sf1 * img.shape[0])), interpolation=random.choice([1, 2, 3]))
img = np.clip(img, 0.0, 1.0)
return img
# def add_Gaussian_noise(img, noise_level1=2, noise_level2=25):
# noise_level = random.randint(noise_level1, noise_level2)
# rnum = np.random.rand()
# if rnum > 0.6: # add color Gaussian noise
# img += np.random.normal(0, noise_level / 255.0, img.shape).astype(np.float32)
# elif rnum < 0.4: # add grayscale Gaussian noise
# img += np.random.normal(0, noise_level / 255.0, (*img.shape[:2], 1)).astype(np.float32)
# else: # add noise
# L = noise_level2 / 255.
# D = np.diag(np.random.rand(3))
# U = orth(np.random.rand(3, 3))
# conv = np.dot(np.dot(np.transpose(U), D), U)
# img += np.random.multivariate_normal([0, 0, 0], np.abs(L ** 2 * conv), img.shape[:2]).astype(np.float32)
# img = np.clip(img, 0.0, 1.0)
# return img
def add_Gaussian_noise(img, noise_level1=2, noise_level2=25):
noise_level = random.randint(noise_level1, noise_level2)
rnum = np.random.rand()
if rnum > 0.6: # add color Gaussian noise
img = img + np.random.normal(0, noise_level / 255.0, img.shape).astype(np.float32)
elif rnum < 0.4: # add grayscale Gaussian noise
img = img + np.random.normal(0, noise_level / 255.0, (*img.shape[:2], 1)).astype(np.float32)
else: # add noise
L = noise_level2 / 255.
D = np.diag(np.random.rand(3))
U = orth(np.random.rand(3, 3))
conv = np.dot(np.dot(np.transpose(U), D), U)
img = img + np.random.multivariate_normal([0, 0, 0], np.abs(L ** 2 * conv), img.shape[:2]).astype(np.float32)
img = np.clip(img, 0.0, 1.0)
return img
def add_speckle_noise(img, noise_level1=2, noise_level2=25):
noise_level = random.randint(noise_level1, noise_level2)
img = np.clip(img, 0.0, 1.0)
rnum = random.random()
if rnum > 0.6:
img += img * np.random.normal(0, noise_level / 255.0, img.shape).astype(np.float32)
elif rnum < 0.4:
img += img * np.random.normal(0, noise_level / 255.0, (*img.shape[:2], 1)).astype(np.float32)
else:
L = noise_level2 / 255.
D = np.diag(np.random.rand(3))
U = orth(np.random.rand(3, 3))
conv = np.dot(np.dot(np.transpose(U), D), U)
img += img * np.random.multivariate_normal([0, 0, 0], np.abs(L ** 2 * conv), img.shape[:2]).astype(np.float32)
img = np.clip(img, 0.0, 1.0)
return img
def add_Poisson_noise(img):
img = np.clip((img * 255.0).round(), 0, 255) / 255.
vals = 10 ** (2 * random.random() + 2.0) # [2, 4]
if random.random() < 0.5:
img = np.random.poisson(img * vals).astype(np.float32) / vals
else:
img_gray = np.dot(img[..., :3], [0.299, 0.587, 0.114])
img_gray = np.clip((img_gray * 255.0).round(), 0, 255) / 255.
noise_gray = np.random.poisson(img_gray * vals).astype(np.float32) / vals - img_gray
img += noise_gray[:, :, np.newaxis]
img = np.clip(img, 0.0, 1.0)
return img
def add_JPEG_noise(img):
quality_factor = random.randint(30, 95)
img = cv2.cvtColor(util.single2uint(img), cv2.COLOR_RGB2BGR)
result, encimg = cv2.imencode('.jpg', img, [int(cv2.IMWRITE_JPEG_QUALITY), quality_factor])
img = cv2.imdecode(encimg, 1)
img = cv2.cvtColor(util.uint2single(img), cv2.COLOR_BGR2RGB)
return img
def random_crop(lq, hq, sf=4, lq_patchsize=64):
h, w = lq.shape[:2]
rnd_h = random.randint(0, h - lq_patchsize)
rnd_w = random.randint(0, w - lq_patchsize)
lq = lq[rnd_h:rnd_h + lq_patchsize, rnd_w:rnd_w + lq_patchsize, :]
rnd_h_H, rnd_w_H = int(rnd_h * sf), int(rnd_w * sf)
hq = hq[rnd_h_H:rnd_h_H + lq_patchsize * sf, rnd_w_H:rnd_w_H + lq_patchsize * sf, :]
return lq, hq
def degradation_bsrgan(img, sf=4, lq_patchsize=72, isp_model=None):
"""
This is the degradation model of BSRGAN from the paper
"Designing a Practical Degradation Model for Deep Blind Image Super-Resolution"
----------
img: HXWXC, [0, 1], its size should be large than (lq_patchsizexsf)x(lq_patchsizexsf)
sf: scale factor
isp_model: camera ISP model
Returns
-------
img: low-quality patch, size: lq_patchsizeXlq_patchsizeXC, range: [0, 1]
hq: corresponding high-quality patch, size: (lq_patchsizexsf)X(lq_patchsizexsf)XC, range: [0, 1]
"""
isp_prob, jpeg_prob, scale2_prob = 0.25, 0.9, 0.25
sf_ori = sf
h1, w1 = img.shape[:2]
img = img.copy()[:w1 - w1 % sf, :h1 - h1 % sf, ...] # mod crop
h, w = img.shape[:2]
if h < lq_patchsize * sf or w < lq_patchsize * sf:
raise ValueError(f'img size ({h1}X{w1}) is too small!')
hq = img.copy()
if sf == 4 and random.random() < scale2_prob: # downsample1
if np.random.rand() < 0.5:
img = cv2.resize(img, (int(1 / 2 * img.shape[1]), int(1 / 2 * img.shape[0])),
interpolation=random.choice([1, 2, 3]))
else:
img = util.imresize_np(img, 1 / 2, True)
img = np.clip(img, 0.0, 1.0)
sf = 2
shuffle_order = random.sample(range(7), 7)
idx1, idx2 = shuffle_order.index(2), shuffle_order.index(3)
if idx1 > idx2: # keep downsample3 last
shuffle_order[idx1], shuffle_order[idx2] = shuffle_order[idx2], shuffle_order[idx1]
for i in shuffle_order:
if i == 0:
img = add_blur(img, sf=sf)
elif i == 1:
img = add_blur(img, sf=sf)
elif i == 2:
a, b = img.shape[1], img.shape[0]
# downsample2
if random.random() < 0.75:
sf1 = random.uniform(1, 2 * sf)
img = cv2.resize(img, (int(1 / sf1 * img.shape[1]), int(1 / sf1 * img.shape[0])),
interpolation=random.choice([1, 2, 3]))
else:
k = fspecial('gaussian', 25, random.uniform(0.1, 0.6 * sf))
k_shifted = shift_pixel(k, sf)
k_shifted = k_shifted / k_shifted.sum() # blur with shifted kernel
img = ndimage.filters.convolve(img, np.expand_dims(k_shifted, axis=2), mode='mirror')
img = img[0::sf, 0::sf, ...] # nearest downsampling
img = np.clip(img, 0.0, 1.0)
elif i == 3:
# downsample3
img = cv2.resize(img, (int(1 / sf * a), int(1 / sf * b)), interpolation=random.choice([1, 2, 3]))
img = np.clip(img, 0.0, 1.0)
elif i == 4:
# add Gaussian noise
img = add_Gaussian_noise(img, noise_level1=2, noise_level2=25)
elif i == 5:
# add JPEG noise
if random.random() < jpeg_prob:
img = add_JPEG_noise(img)
elif i == 6:
# add processed camera sensor noise
if random.random() < isp_prob and isp_model is not None:
with torch.no_grad():
img, hq = isp_model.forward(img.copy(), hq)
# add final JPEG compression noise
img = add_JPEG_noise(img)
# random crop
img, hq = random_crop(img, hq, sf_ori, lq_patchsize)
return img, hq
# todo no isp_model?
def degradation_bsrgan_variant(image, sf=4, isp_model=None):
"""
This is the degradation model of BSRGAN from the paper
"Designing a Practical Degradation Model for Deep Blind Image Super-Resolution"
----------
sf: scale factor
isp_model: camera ISP model
Returns
-------
img: low-quality patch, size: lq_patchsizeXlq_patchsizeXC, range: [0, 1]
hq: corresponding high-quality patch, size: (lq_patchsizexsf)X(lq_patchsizexsf)XC, range: [0, 1]
"""
image = util.uint2single(image)
isp_prob, jpeg_prob, scale2_prob = 0.25, 0.9, 0.25
sf_ori = sf
h1, w1 = image.shape[:2]
image = image.copy()[:w1 - w1 % sf, :h1 - h1 % sf, ...] # mod crop
h, w = image.shape[:2]
hq = image.copy()
if sf == 4 and random.random() < scale2_prob: # downsample1
if np.random.rand() < 0.5:
image = cv2.resize(image, (int(1 / 2 * image.shape[1]), int(1 / 2 * image.shape[0])),
interpolation=random.choice([1, 2, 3]))
else:
image = util.imresize_np(image, 1 / 2, True)
image = np.clip(image, 0.0, 1.0)
sf = 2
shuffle_order = random.sample(range(7), 7)
idx1, idx2 = shuffle_order.index(2), shuffle_order.index(3)
if idx1 > idx2: # keep downsample3 last
shuffle_order[idx1], shuffle_order[idx2] = shuffle_order[idx2], shuffle_order[idx1]
for i in shuffle_order:
if i == 0:
image = add_blur(image, sf=sf)
elif i == 1:
image = add_blur(image, sf=sf)
elif i == 2:
a, b = image.shape[1], image.shape[0]
# downsample2
if random.random() < 0.75:
sf1 = random.uniform(1, 2 * sf)
image = cv2.resize(image, (int(1 / sf1 * image.shape[1]), int(1 / sf1 * image.shape[0])),
interpolation=random.choice([1, 2, 3]))
else:
k = fspecial('gaussian', 25, random.uniform(0.1, 0.6 * sf))
k_shifted = shift_pixel(k, sf)
k_shifted = k_shifted / k_shifted.sum() # blur with shifted kernel
image = ndimage.filters.convolve(image, np.expand_dims(k_shifted, axis=2), mode='mirror')
image = image[0::sf, 0::sf, ...] # nearest downsampling
image = np.clip(image, 0.0, 1.0)
elif i == 3:
# downsample3
image = cv2.resize(image, (int(1 / sf * a), int(1 / sf * b)), interpolation=random.choice([1, 2, 3]))
image = np.clip(image, 0.0, 1.0)
elif i == 4:
# add Gaussian noise
image = add_Gaussian_noise(image, noise_level1=2, noise_level2=25)
elif i == 5:
# add JPEG noise
if random.random() < jpeg_prob:
image = add_JPEG_noise(image)
# elif i == 6:
# # add processed camera sensor noise
# if random.random() < isp_prob and isp_model is not None:
# with torch.no_grad():
# img, hq = isp_model.forward(img.copy(), hq)
# add final JPEG compression noise
image = add_JPEG_noise(image)
image = util.single2uint(image)
example = {"image":image}
return example
# TODO incase there is a pickle error one needs to replace a += x with a = a + x in add_speckle_noise etc...
def degradation_bsrgan_plus(img, sf=4, shuffle_prob=0.5, use_sharp=True, lq_patchsize=64, isp_model=None):
"""
This is an extended degradation model by combining
the degradation models of BSRGAN and Real-ESRGAN
----------
img: HXWXC, [0, 1], its size should be large than (lq_patchsizexsf)x(lq_patchsizexsf)
sf: scale factor
use_shuffle: the degradation shuffle
use_sharp: sharpening the img
Returns
-------
img: low-quality patch, size: lq_patchsizeXlq_patchsizeXC, range: [0, 1]
hq: corresponding high-quality patch, size: (lq_patchsizexsf)X(lq_patchsizexsf)XC, range: [0, 1]
"""
h1, w1 = img.shape[:2]
img = img.copy()[:w1 - w1 % sf, :h1 - h1 % sf, ...] # mod crop
h, w = img.shape[:2]
if h < lq_patchsize * sf or w < lq_patchsize * sf:
raise ValueError(f'img size ({h1}X{w1}) is too small!')
if use_sharp:
img = add_sharpening(img)
hq = img.copy()
if random.random() < shuffle_prob:
shuffle_order = random.sample(range(13), 13)
else:
shuffle_order = list(range(13))
# local shuffle for noise, JPEG is always the last one
shuffle_order[2:6] = random.sample(shuffle_order[2:6], len(range(2, 6)))
shuffle_order[9:13] = random.sample(shuffle_order[9:13], len(range(9, 13)))
poisson_prob, speckle_prob, isp_prob = 0.1, 0.1, 0.1
for i in shuffle_order:
if i == 0:
img = add_blur(img, sf=sf)
elif i == 1:
img = add_resize(img, sf=sf)
elif i == 2:
img = add_Gaussian_noise(img, noise_level1=2, noise_level2=25)
elif i == 3:
if random.random() < poisson_prob:
img = add_Poisson_noise(img)
elif i == 4:
if random.random() < speckle_prob:
img = add_speckle_noise(img)
elif i == 5:
if random.random() < isp_prob and isp_model is not None:
with torch.no_grad():
img, hq = isp_model.forward(img.copy(), hq)
elif i == 6:
img = add_JPEG_noise(img)
elif i == 7:
img = add_blur(img, sf=sf)
elif i == 8:
img = add_resize(img, sf=sf)
elif i == 9:
img = add_Gaussian_noise(img, noise_level1=2, noise_level2=25)
elif i == 10:
if random.random() < poisson_prob:
img = add_Poisson_noise(img)
elif i == 11:
if random.random() < speckle_prob:
img = add_speckle_noise(img)
elif i == 12:
if random.random() < isp_prob and isp_model is not None:
with torch.no_grad():
img, hq = isp_model.forward(img.copy(), hq)
else:
print('check the shuffle!')
# resize to desired size
img = cv2.resize(img, (int(1 / sf * hq.shape[1]), int(1 / sf * hq.shape[0])),
interpolation=random.choice([1, 2, 3]))
# add final JPEG compression noise
img = add_JPEG_noise(img)
# random crop
img, hq = random_crop(img, hq, sf, lq_patchsize)
return img, hq
if __name__ == '__main__':
print("hey")
img = util.imread_uint('utils/test.png', 3)
print(img)
img = util.uint2single(img)
print(img)
img = img[:448, :448]
h = img.shape[0] // 4
print("resizing to", h)
sf = 4
deg_fn = partial(degradation_bsrgan_variant, sf=sf)
for i in range(20):
print(i)
img_lq = deg_fn(img)
print(img_lq)
img_lq_bicubic = albumentations.SmallestMaxSize(max_size=h, interpolation=cv2.INTER_CUBIC)(image=img)["image"]
print(img_lq.shape)
print("bicubic", img_lq_bicubic.shape)
print(img_hq.shape)
lq_nearest = cv2.resize(util.single2uint(img_lq), (int(sf * img_lq.shape[1]), int(sf * img_lq.shape[0])),
interpolation=0)
lq_bicubic_nearest = cv2.resize(util.single2uint(img_lq_bicubic), (int(sf * img_lq.shape[1]), int(sf * img_lq.shape[0])),
interpolation=0)
img_concat = np.concatenate([lq_bicubic_nearest, lq_nearest, util.single2uint(img_hq)], axis=1)
util.imsave(img_concat, str(i) + '.png')
@@ -0,0 +1,650 @@
# -*- coding: utf-8 -*-
import random
from functools import partial
import albumentations
import cv2
import numpy as np
import scipy
import scipy.stats as ss
import torch
from scipy import ndimage
from scipy.interpolate import interp2d
from scipy.linalg import orth
from . import utils_image as util
"""
# --------------------------------------------
# Super-Resolution
# --------------------------------------------
#
# Kai Zhang (cskaizhang@gmail.com)
# https://github.com/cszn
# From 2019/03--2021/08
# --------------------------------------------
"""
def modcrop_np(img, sf):
'''
Args:
img: numpy image, WxH or WxHxC
sf: scale factor
Return:
cropped image
'''
w, h = img.shape[:2]
im = np.copy(img)
return im[:w - w % sf, :h - h % sf, ...]
"""
# --------------------------------------------
# anisotropic Gaussian kernels
# --------------------------------------------
"""
def analytic_kernel(k):
"""Calculate the X4 kernel from the X2 kernel (for proof see appendix in paper)"""
k_size = k.shape[0]
# Calculate the big kernels size
big_k = np.zeros((3 * k_size - 2, 3 * k_size - 2))
# Loop over the small kernel to fill the big one
for r in range(k_size):
for c in range(k_size):
big_k[2 * r:2 * r + k_size, 2 * c:2 * c + k_size] += k[r, c] * k
# Crop the edges of the big kernel to ignore very small values and increase run time of SR
crop = k_size // 2
cropped_big_k = big_k[crop:-crop, crop:-crop]
# Normalize to 1
return cropped_big_k / cropped_big_k.sum()
def anisotropic_Gaussian(ksize=15, theta=np.pi, l1=6, l2=6):
""" generate an anisotropic Gaussian kernel
Args:
ksize : e.g., 15, kernel size
theta : [0, pi], rotation angle range
l1 : [0.1,50], scaling of eigenvalues
l2 : [0.1,l1], scaling of eigenvalues
If l1 = l2, will get an isotropic Gaussian kernel.
Returns:
k : kernel
"""
v = np.dot(np.array([[np.cos(theta), -np.sin(theta)], [np.sin(theta), np.cos(theta)]]), np.array([1., 0.]))
V = np.array([[v[0], v[1]], [v[1], -v[0]]])
D = np.array([[l1, 0], [0, l2]])
Sigma = np.dot(np.dot(V, D), np.linalg.inv(V))
k = gm_blur_kernel(mean=[0, 0], cov=Sigma, size=ksize)
return k
def gm_blur_kernel(mean, cov, size=15):
center = size / 2.0 + 0.5
k = np.zeros([size, size])
for y in range(size):
for x in range(size):
cy = y - center + 1
cx = x - center + 1
k[y, x] = ss.multivariate_normal.pdf([cx, cy], mean=mean, cov=cov)
k = k / np.sum(k)
return k
def shift_pixel(x, sf, upper_left=True):
"""shift pixel for super-resolution with different scale factors
Args:
x: WxHxC or WxH
sf: scale factor
upper_left: shift direction
"""
h, w = x.shape[:2]
shift = (sf - 1) * 0.5
xv, yv = np.arange(0, w, 1.0), np.arange(0, h, 1.0)
if upper_left:
x1 = xv + shift
y1 = yv + shift
else:
x1 = xv - shift
y1 = yv - shift
x1 = np.clip(x1, 0, w - 1)
y1 = np.clip(y1, 0, h - 1)
if x.ndim == 2:
x = interp2d(xv, yv, x)(x1, y1)
if x.ndim == 3:
for i in range(x.shape[-1]):
x[:, :, i] = interp2d(xv, yv, x[:, :, i])(x1, y1)
return x
def blur(x, k):
'''
x: image, NxcxHxW
k: kernel, Nx1xhxw
'''
n, c = x.shape[:2]
p1, p2 = (k.shape[-2] - 1) // 2, (k.shape[-1] - 1) // 2
x = torch.nn.functional.pad(x, pad=(p1, p2, p1, p2), mode='replicate')
k = k.repeat(1, c, 1, 1)
k = k.view(-1, 1, k.shape[2], k.shape[3])
x = x.view(1, -1, x.shape[2], x.shape[3])
x = torch.nn.functional.conv2d(x, k, bias=None, stride=1, padding=0, groups=n * c)
x = x.view(n, c, x.shape[2], x.shape[3])
return x
def gen_kernel(k_size=np.array([15, 15]), scale_factor=np.array([4, 4]), min_var=0.6, max_var=10., noise_level=0):
""""
# modified version of https://github.com/assafshocher/BlindSR_dataset_generator
# Kai Zhang
# min_var = 0.175 * sf # variance of the gaussian kernel will be sampled between min_var and max_var
# max_var = 2.5 * sf
"""
# Set random eigen-vals (lambdas) and angle (theta) for COV matrix
lambda_1 = min_var + np.random.rand() * (max_var - min_var)
lambda_2 = min_var + np.random.rand() * (max_var - min_var)
theta = np.random.rand() * np.pi # random theta
noise = -noise_level + np.random.rand(*k_size) * noise_level * 2
# Set COV matrix using Lambdas and Theta
LAMBDA = np.diag([lambda_1, lambda_2])
Q = np.array([[np.cos(theta), -np.sin(theta)],
[np.sin(theta), np.cos(theta)]])
SIGMA = Q @ LAMBDA @ Q.T
INV_SIGMA = np.linalg.inv(SIGMA)[None, None, :, :]
# Set expectation position (shifting kernel for aligned image)
MU = k_size // 2 - 0.5 * (scale_factor - 1) # - 0.5 * (scale_factor - k_size % 2)
MU = MU[None, None, :, None]
# Create meshgrid for Gaussian
[X, Y] = np.meshgrid(range(k_size[0]), range(k_size[1]))
Z = np.stack([X, Y], 2)[:, :, :, None]
# Calcualte Gaussian for every pixel of the kernel
ZZ = Z - MU
ZZ_t = ZZ.transpose(0, 1, 3, 2)
raw_kernel = np.exp(-0.5 * np.squeeze(ZZ_t @ INV_SIGMA @ ZZ)) * (1 + noise)
# shift the kernel so it will be centered
# raw_kernel_centered = kernel_shift(raw_kernel, scale_factor)
# Normalize the kernel and return
# kernel = raw_kernel_centered / np.sum(raw_kernel_centered)
kernel = raw_kernel / np.sum(raw_kernel)
return kernel
def fspecial_gaussian(hsize, sigma):
hsize = [hsize, hsize]
siz = [(hsize[0] - 1.0) / 2.0, (hsize[1] - 1.0) / 2.0]
std = sigma
[x, y] = np.meshgrid(np.arange(-siz[1], siz[1] + 1), np.arange(-siz[0], siz[0] + 1))
arg = -(x * x + y * y) / (2 * std * std)
h = np.exp(arg)
h[h < scipy.finfo(float).eps * h.max()] = 0
sumh = h.sum()
if sumh != 0:
h = h / sumh
return h
def fspecial_laplacian(alpha):
alpha = max([0, min([alpha, 1])])
h1 = alpha / (alpha + 1)
h2 = (1 - alpha) / (alpha + 1)
h = [[h1, h2, h1], [h2, -4 / (alpha + 1), h2], [h1, h2, h1]]
h = np.array(h)
return h
def fspecial(filter_type, *args, **kwargs):
'''
python code from:
https://github.com/ronaldosena/imagens-medicas-2/blob/40171a6c259edec7827a6693a93955de2bd39e76/Aulas/aula_2_-_uniform_filter/matlab_fspecial.py
'''
if filter_type == 'gaussian':
return fspecial_gaussian(*args, **kwargs)
if filter_type == 'laplacian':
return fspecial_laplacian(*args, **kwargs)
"""
# --------------------------------------------
# degradation models
# --------------------------------------------
"""
def bicubic_degradation(x, sf=3):
'''
Args:
x: HxWxC image, [0, 1]
sf: down-scale factor
Return:
bicubicly downsampled LR image
'''
x = util.imresize_np(x, scale=1 / sf)
return x
def srmd_degradation(x, k, sf=3):
''' blur + bicubic downsampling
Args:
x: HxWxC image, [0, 1]
k: hxw, double
sf: down-scale factor
Return:
downsampled LR image
Reference:
@inproceedings{zhang2018learning,
title={Learning a single convolutional super-resolution network for multiple degradations},
author={Zhang, Kai and Zuo, Wangmeng and Zhang, Lei},
booktitle={IEEE Conference on Computer Vision and Pattern Recognition},
pages={3262--3271},
year={2018}
}
'''
x = ndimage.filters.convolve(x, np.expand_dims(k, axis=2), mode='wrap') # 'nearest' | 'mirror'
x = bicubic_degradation(x, sf=sf)
return x
def dpsr_degradation(x, k, sf=3):
''' bicubic downsampling + blur
Args:
x: HxWxC image, [0, 1]
k: hxw, double
sf: down-scale factor
Return:
downsampled LR image
Reference:
@inproceedings{zhang2019deep,
title={Deep Plug-and-Play Super-Resolution for Arbitrary Blur Kernels},
author={Zhang, Kai and Zuo, Wangmeng and Zhang, Lei},
booktitle={IEEE Conference on Computer Vision and Pattern Recognition},
pages={1671--1681},
year={2019}
}
'''
x = bicubic_degradation(x, sf=sf)
x = ndimage.filters.convolve(x, np.expand_dims(k, axis=2), mode='wrap')
return x
def classical_degradation(x, k, sf=3):
''' blur + downsampling
Args:
x: HxWxC image, [0, 1]/[0, 255]
k: hxw, double
sf: down-scale factor
Return:
downsampled LR image
'''
x = ndimage.filters.convolve(x, np.expand_dims(k, axis=2), mode='wrap')
# x = filters.correlate(x, np.expand_dims(np.flip(k), axis=2))
st = 0
return x[st::sf, st::sf, ...]
def add_sharpening(img, weight=0.5, radius=50, threshold=10):
"""USM sharpening. borrowed from real-ESRGAN
Input image: I; Blurry image: B.
1. K = I + weight * (I - B)
2. Mask = 1 if abs(I - B) > threshold, else: 0
3. Blur mask:
4. Out = Mask * K + (1 - Mask) * I
Args:
img (Numpy array): Input image, HWC, BGR; float32, [0, 1].
weight (float): Sharp weight. Default: 1.
radius (float): Kernel size of Gaussian blur. Default: 50.
threshold (int):
"""
if radius % 2 == 0:
radius += 1
blur = cv2.GaussianBlur(img, (radius, radius), 0)
residual = img - blur
mask = np.abs(residual) * 255 > threshold
mask = mask.astype('float32')
soft_mask = cv2.GaussianBlur(mask, (radius, radius), 0)
K = img + weight * residual
K = np.clip(K, 0, 1)
return soft_mask * K + (1 - soft_mask) * img
def add_blur(img, sf=4):
wd2 = 4.0 + sf
wd = 2.0 + 0.2 * sf
wd2 = wd2/4
wd = wd/4
if random.random() < 0.5:
l1 = wd2 * random.random()
l2 = wd2 * random.random()
k = anisotropic_Gaussian(ksize=random.randint(2, 11) + 3, theta=random.random() * np.pi, l1=l1, l2=l2)
else:
k = fspecial('gaussian', random.randint(2, 4) + 3, wd * random.random())
img = ndimage.filters.convolve(img, np.expand_dims(k, axis=2), mode='mirror')
return img
def add_resize(img, sf=4):
rnum = np.random.rand()
if rnum > 0.8: # up
sf1 = random.uniform(1, 2)
elif rnum < 0.7: # down
sf1 = random.uniform(0.5 / sf, 1)
else:
sf1 = 1.0
img = cv2.resize(img, (int(sf1 * img.shape[1]), int(sf1 * img.shape[0])), interpolation=random.choice([1, 2, 3]))
img = np.clip(img, 0.0, 1.0)
return img
# def add_Gaussian_noise(img, noise_level1=2, noise_level2=25):
# noise_level = random.randint(noise_level1, noise_level2)
# rnum = np.random.rand()
# if rnum > 0.6: # add color Gaussian noise
# img += np.random.normal(0, noise_level / 255.0, img.shape).astype(np.float32)
# elif rnum < 0.4: # add grayscale Gaussian noise
# img += np.random.normal(0, noise_level / 255.0, (*img.shape[:2], 1)).astype(np.float32)
# else: # add noise
# L = noise_level2 / 255.
# D = np.diag(np.random.rand(3))
# U = orth(np.random.rand(3, 3))
# conv = np.dot(np.dot(np.transpose(U), D), U)
# img += np.random.multivariate_normal([0, 0, 0], np.abs(L ** 2 * conv), img.shape[:2]).astype(np.float32)
# img = np.clip(img, 0.0, 1.0)
# return img
def add_Gaussian_noise(img, noise_level1=2, noise_level2=25):
noise_level = random.randint(noise_level1, noise_level2)
rnum = np.random.rand()
if rnum > 0.6: # add color Gaussian noise
img = img + np.random.normal(0, noise_level / 255.0, img.shape).astype(np.float32)
elif rnum < 0.4: # add grayscale Gaussian noise
img = img + np.random.normal(0, noise_level / 255.0, (*img.shape[:2], 1)).astype(np.float32)
else: # add noise
L = noise_level2 / 255.
D = np.diag(np.random.rand(3))
U = orth(np.random.rand(3, 3))
conv = np.dot(np.dot(np.transpose(U), D), U)
img = img + np.random.multivariate_normal([0, 0, 0], np.abs(L ** 2 * conv), img.shape[:2]).astype(np.float32)
img = np.clip(img, 0.0, 1.0)
return img
def add_speckle_noise(img, noise_level1=2, noise_level2=25):
noise_level = random.randint(noise_level1, noise_level2)
img = np.clip(img, 0.0, 1.0)
rnum = random.random()
if rnum > 0.6:
img += img * np.random.normal(0, noise_level / 255.0, img.shape).astype(np.float32)
elif rnum < 0.4:
img += img * np.random.normal(0, noise_level / 255.0, (*img.shape[:2], 1)).astype(np.float32)
else:
L = noise_level2 / 255.
D = np.diag(np.random.rand(3))
U = orth(np.random.rand(3, 3))
conv = np.dot(np.dot(np.transpose(U), D), U)
img += img * np.random.multivariate_normal([0, 0, 0], np.abs(L ** 2 * conv), img.shape[:2]).astype(np.float32)
img = np.clip(img, 0.0, 1.0)
return img
def add_Poisson_noise(img):
img = np.clip((img * 255.0).round(), 0, 255) / 255.
vals = 10 ** (2 * random.random() + 2.0) # [2, 4]
if random.random() < 0.5:
img = np.random.poisson(img * vals).astype(np.float32) / vals
else:
img_gray = np.dot(img[..., :3], [0.299, 0.587, 0.114])
img_gray = np.clip((img_gray * 255.0).round(), 0, 255) / 255.
noise_gray = np.random.poisson(img_gray * vals).astype(np.float32) / vals - img_gray
img += noise_gray[:, :, np.newaxis]
img = np.clip(img, 0.0, 1.0)
return img
def add_JPEG_noise(img):
quality_factor = random.randint(80, 95)
img = cv2.cvtColor(util.single2uint(img), cv2.COLOR_RGB2BGR)
result, encimg = cv2.imencode('.jpg', img, [int(cv2.IMWRITE_JPEG_QUALITY), quality_factor])
img = cv2.imdecode(encimg, 1)
img = cv2.cvtColor(util.uint2single(img), cv2.COLOR_BGR2RGB)
return img
def random_crop(lq, hq, sf=4, lq_patchsize=64):
h, w = lq.shape[:2]
rnd_h = random.randint(0, h - lq_patchsize)
rnd_w = random.randint(0, w - lq_patchsize)
lq = lq[rnd_h:rnd_h + lq_patchsize, rnd_w:rnd_w + lq_patchsize, :]
rnd_h_H, rnd_w_H = int(rnd_h * sf), int(rnd_w * sf)
hq = hq[rnd_h_H:rnd_h_H + lq_patchsize * sf, rnd_w_H:rnd_w_H + lq_patchsize * sf, :]
return lq, hq
def degradation_bsrgan(img, sf=4, lq_patchsize=72, isp_model=None):
"""
This is the degradation model of BSRGAN from the paper
"Designing a Practical Degradation Model for Deep Blind Image Super-Resolution"
----------
img: HXWXC, [0, 1], its size should be large than (lq_patchsizexsf)x(lq_patchsizexsf)
sf: scale factor
isp_model: camera ISP model
Returns
-------
img: low-quality patch, size: lq_patchsizeXlq_patchsizeXC, range: [0, 1]
hq: corresponding high-quality patch, size: (lq_patchsizexsf)X(lq_patchsizexsf)XC, range: [0, 1]
"""
isp_prob, jpeg_prob, scale2_prob = 0.25, 0.9, 0.25
sf_ori = sf
h1, w1 = img.shape[:2]
img = img.copy()[:w1 - w1 % sf, :h1 - h1 % sf, ...] # mod crop
h, w = img.shape[:2]
if h < lq_patchsize * sf or w < lq_patchsize * sf:
raise ValueError(f'img size ({h1}X{w1}) is too small!')
hq = img.copy()
if sf == 4 and random.random() < scale2_prob: # downsample1
if np.random.rand() < 0.5:
img = cv2.resize(img, (int(1 / 2 * img.shape[1]), int(1 / 2 * img.shape[0])),
interpolation=random.choice([1, 2, 3]))
else:
img = util.imresize_np(img, 1 / 2, True)
img = np.clip(img, 0.0, 1.0)
sf = 2
shuffle_order = random.sample(range(7), 7)
idx1, idx2 = shuffle_order.index(2), shuffle_order.index(3)
if idx1 > idx2: # keep downsample3 last
shuffle_order[idx1], shuffle_order[idx2] = shuffle_order[idx2], shuffle_order[idx1]
for i in shuffle_order:
if i == 0:
img = add_blur(img, sf=sf)
elif i == 1:
img = add_blur(img, sf=sf)
elif i == 2:
a, b = img.shape[1], img.shape[0]
# downsample2
if random.random() < 0.75:
sf1 = random.uniform(1, 2 * sf)
img = cv2.resize(img, (int(1 / sf1 * img.shape[1]), int(1 / sf1 * img.shape[0])),
interpolation=random.choice([1, 2, 3]))
else:
k = fspecial('gaussian', 25, random.uniform(0.1, 0.6 * sf))
k_shifted = shift_pixel(k, sf)
k_shifted = k_shifted / k_shifted.sum() # blur with shifted kernel
img = ndimage.filters.convolve(img, np.expand_dims(k_shifted, axis=2), mode='mirror')
img = img[0::sf, 0::sf, ...] # nearest downsampling
img = np.clip(img, 0.0, 1.0)
elif i == 3:
# downsample3
img = cv2.resize(img, (int(1 / sf * a), int(1 / sf * b)), interpolation=random.choice([1, 2, 3]))
img = np.clip(img, 0.0, 1.0)
elif i == 4:
# add Gaussian noise
img = add_Gaussian_noise(img, noise_level1=2, noise_level2=8)
elif i == 5:
# add JPEG noise
if random.random() < jpeg_prob:
img = add_JPEG_noise(img)
elif i == 6:
# add processed camera sensor noise
if random.random() < isp_prob and isp_model is not None:
with torch.no_grad():
img, hq = isp_model.forward(img.copy(), hq)
# add final JPEG compression noise
img = add_JPEG_noise(img)
# random crop
img, hq = random_crop(img, hq, sf_ori, lq_patchsize)
return img, hq
# todo no isp_model?
def degradation_bsrgan_variant(image, sf=4, isp_model=None):
"""
This is the degradation model of BSRGAN from the paper
"Designing a Practical Degradation Model for Deep Blind Image Super-Resolution"
----------
sf: scale factor
isp_model: camera ISP model
Returns
-------
img: low-quality patch, size: lq_patchsizeXlq_patchsizeXC, range: [0, 1]
hq: corresponding high-quality patch, size: (lq_patchsizexsf)X(lq_patchsizexsf)XC, range: [0, 1]
"""
image = util.uint2single(image)
isp_prob, jpeg_prob, scale2_prob = 0.25, 0.9, 0.25
sf_ori = sf
h1, w1 = image.shape[:2]
image = image.copy()[:w1 - w1 % sf, :h1 - h1 % sf, ...] # mod crop
h, w = image.shape[:2]
hq = image.copy()
if sf == 4 and random.random() < scale2_prob: # downsample1
if np.random.rand() < 0.5:
image = cv2.resize(image, (int(1 / 2 * image.shape[1]), int(1 / 2 * image.shape[0])),
interpolation=random.choice([1, 2, 3]))
else:
image = util.imresize_np(image, 1 / 2, True)
image = np.clip(image, 0.0, 1.0)
sf = 2
shuffle_order = random.sample(range(7), 7)
idx1, idx2 = shuffle_order.index(2), shuffle_order.index(3)
if idx1 > idx2: # keep downsample3 last
shuffle_order[idx1], shuffle_order[idx2] = shuffle_order[idx2], shuffle_order[idx1]
for i in shuffle_order:
if i == 0:
image = add_blur(image, sf=sf)
# elif i == 1:
# image = add_blur(image, sf=sf)
if i == 0:
pass
elif i == 2:
a, b = image.shape[1], image.shape[0]
# downsample2
if random.random() < 0.8:
sf1 = random.uniform(1, 2 * sf)
image = cv2.resize(image, (int(1 / sf1 * image.shape[1]), int(1 / sf1 * image.shape[0])),
interpolation=random.choice([1, 2, 3]))
else:
k = fspecial('gaussian', 25, random.uniform(0.1, 0.6 * sf))
k_shifted = shift_pixel(k, sf)
k_shifted = k_shifted / k_shifted.sum() # blur with shifted kernel
image = ndimage.filters.convolve(image, np.expand_dims(k_shifted, axis=2), mode='mirror')
image = image[0::sf, 0::sf, ...] # nearest downsampling
image = np.clip(image, 0.0, 1.0)
elif i == 3:
# downsample3
image = cv2.resize(image, (int(1 / sf * a), int(1 / sf * b)), interpolation=random.choice([1, 2, 3]))
image = np.clip(image, 0.0, 1.0)
elif i == 4:
# add Gaussian noise
image = add_Gaussian_noise(image, noise_level1=1, noise_level2=2)
elif i == 5:
# add JPEG noise
if random.random() < jpeg_prob:
image = add_JPEG_noise(image)
#
# elif i == 6:
# # add processed camera sensor noise
# if random.random() < isp_prob and isp_model is not None:
# with torch.no_grad():
# img, hq = isp_model.forward(img.copy(), hq)
# add final JPEG compression noise
image = add_JPEG_noise(image)
image = util.single2uint(image)
example = {"image": image}
return example
if __name__ == '__main__':
print("hey")
img = util.imread_uint('utils/test.png', 3)
img = img[:448, :448]
h = img.shape[0] // 4
print("resizing to", h)
sf = 4
deg_fn = partial(degradation_bsrgan_variant, sf=sf)
for i in range(20):
print(i)
img_hq = img
img_lq = deg_fn(img)["image"]
img_hq, img_lq = util.uint2single(img_hq), util.uint2single(img_lq)
print(img_lq)
img_lq_bicubic = albumentations.SmallestMaxSize(max_size=h, interpolation=cv2.INTER_CUBIC)(image=img_hq)["image"]
print(img_lq.shape)
print("bicubic", img_lq_bicubic.shape)
print(img_hq.shape)
lq_nearest = cv2.resize(util.single2uint(img_lq), (int(sf * img_lq.shape[1]), int(sf * img_lq.shape[0])),
interpolation=0)
lq_bicubic_nearest = cv2.resize(util.single2uint(img_lq_bicubic),
(int(sf * img_lq.shape[1]), int(sf * img_lq.shape[0])),
interpolation=0)
img_concat = np.concatenate([lq_bicubic_nearest, lq_nearest, util.single2uint(img_hq)], axis=1)
util.imsave(img_concat, str(i) + '.png')

Some files were not shown because too many files have changed in this diff Show More