initial commit
This commit is contained in:
Executable
+171
@@ -0,0 +1,171 @@
|
||||
outputs/
|
||||
processed/
|
||||
profile/
|
||||
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
|
||||
# C extensions
|
||||
*.so
|
||||
|
||||
# Distribution / packaging
|
||||
.Python
|
||||
build/
|
||||
develop-eggs/
|
||||
dist/
|
||||
downloads/
|
||||
eggs/
|
||||
.eggs/
|
||||
lib/
|
||||
lib64/
|
||||
parts/
|
||||
sdist/
|
||||
var/
|
||||
wheels/
|
||||
pip-wheel-metadata/
|
||||
share/python-wheels/
|
||||
*.egg-info/
|
||||
.installed.cfg
|
||||
*.egg
|
||||
MANIFEST
|
||||
|
||||
# PyInstaller
|
||||
# Usually these files are written by a python script from a template
|
||||
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
||||
*.manifest
|
||||
*.spec
|
||||
|
||||
# Installer logs
|
||||
pip-log.txt
|
||||
pip-delete-this-directory.txt
|
||||
|
||||
# Unit test / coverage reports
|
||||
htmlcov/
|
||||
.tox/
|
||||
.nox/
|
||||
.coverage
|
||||
.coverage.*
|
||||
.cache
|
||||
nosetests.xml
|
||||
coverage.xml
|
||||
*.cover
|
||||
*.py,cover
|
||||
.hypothesis/
|
||||
.pytest_cache/
|
||||
|
||||
# Translations
|
||||
*.mo
|
||||
*.pot
|
||||
|
||||
# Django stuff:
|
||||
*.log
|
||||
local_settings.py
|
||||
db.sqlite3
|
||||
db.sqlite3-journal
|
||||
|
||||
# Flask stuff:
|
||||
instance/
|
||||
.webassets-cache
|
||||
|
||||
# Scrapy stuff:
|
||||
.scrapy
|
||||
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
docs/.build/
|
||||
|
||||
# PyBuilder
|
||||
target/
|
||||
|
||||
# Jupyter Notebook
|
||||
.ipynb_checkpoints
|
||||
|
||||
# IPython
|
||||
profile_default/
|
||||
ipython_config.py
|
||||
|
||||
# pyenv
|
||||
.python-version
|
||||
|
||||
# pipenv
|
||||
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
||||
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
||||
# having no cross-platform support, pipenv may install dependencies that don't work, or not
|
||||
# install all needed dependencies.
|
||||
#Pipfile.lock
|
||||
|
||||
# PEP 582; used by e.g. github.com/David-OConnor/pyflow
|
||||
__pypackages__/
|
||||
|
||||
# Celery stuff
|
||||
celerybeat-schedule
|
||||
celerybeat.pid
|
||||
|
||||
# SageMath parsed files
|
||||
*.sage.py
|
||||
|
||||
# Environments
|
||||
.env
|
||||
.venv
|
||||
env/
|
||||
venv/
|
||||
ENV/
|
||||
env.bak/
|
||||
venv.bak/
|
||||
|
||||
# Spyder project settings
|
||||
.spyderproject
|
||||
.spyproject
|
||||
|
||||
# Rope project settings
|
||||
.ropeproject
|
||||
|
||||
# mkdocs documentation
|
||||
/site
|
||||
|
||||
# mypy
|
||||
.mypy_cache/
|
||||
.dmypy.json
|
||||
dmypy.json
|
||||
|
||||
# Pyre type checker
|
||||
.pyre/
|
||||
|
||||
# IDE
|
||||
.idea/
|
||||
.vscode/
|
||||
|
||||
# macos
|
||||
*.DS_Store
|
||||
#data/
|
||||
|
||||
docs/.build
|
||||
|
||||
# pytorch checkpoint
|
||||
*.pt
|
||||
|
||||
# ignore version.py generated by setup.py
|
||||
colossalai/version.py
|
||||
|
||||
# ignore any kernel build files
|
||||
.o
|
||||
.so
|
||||
|
||||
# ignore python interface defition file
|
||||
.pyi
|
||||
|
||||
# ignore coverage test file
|
||||
coverage.lcov
|
||||
coverage.xml
|
||||
|
||||
# ignore testmon and coverage files
|
||||
.coverage
|
||||
.testmondata*
|
||||
|
||||
pretrained
|
||||
samples
|
||||
cache_dir
|
||||
taming
|
||||
evaluations/pab/datasets/
|
||||
@@ -0,0 +1,150 @@
|
||||
<p align="center">
|
||||
<img width="200px" alt="OpenDiT" src="./assets/figures/logo.png?raw=true">
|
||||
</p>
|
||||
<p align="center"><b><big>An Easy, Fast and Memory-Efficient System for DiT Training and Inference</big></b></p>
|
||||
</p>
|
||||
|
||||
### Latest News 🔥
|
||||
- [2024/06] 🔥<b>Propose Pyramid Attention Broadcast (PAB)[[blog](https://oahzxl.github.io/PAB/)][[doc](./docs/pab.md)], the first approach to achieve <b>real-time</b> DiT-based video generation, delivering <b>negligible quality loss</b> without <b>requiring any training</b>.</b>
|
||||
- [2024/06] Support Open-Sora-Plan and Latte.
|
||||
- [2024/03] Propose Dynamic Sequence Parallel (DSP)[[paper](https://arxiv.org/abs/2403.10266)][[doc](./docs/dsp.md)], achieves **3x** speed for training and **2x** speed for inference in OpenSora compared with sota sequence parallelism.
|
||||
- [2024/03] Support Open-Sora: Democratizing Efficient Video Production for All.
|
||||
- [2024/02] Release OpenDiT: An Easy, Fast and Memory-Efficent System for DiT Training and Inference.
|
||||
|
||||
# About
|
||||
|
||||
OpenDiT is an open-source project that provides a high-performance implementation of Diffusion Transformer (DiT) powered by Colossal-AI, specifically designed to enhance the efficiency of training and inference for DiT applications, including text-to-video generation and text-to-image generation.
|
||||
|
||||
OpenDiT will continue to integrate more open-source DiT models and techniques. Stay tuned for upcoming enhancements and additional features!
|
||||
|
||||
## Installation
|
||||
|
||||
Prerequisites:
|
||||
|
||||
- Python >= 3.10
|
||||
- PyTorch >= 1.13 (We recommend to use a >2.0 version)
|
||||
- CUDA >= 11.6
|
||||
|
||||
We strongly recommend using Anaconda to create a new environment (Python >= 3.10) to run our examples:
|
||||
|
||||
```shell
|
||||
conda create -n opendit python=3.10 -y
|
||||
conda activate opendit
|
||||
```
|
||||
|
||||
Install ColossalAI:
|
||||
|
||||
```shell
|
||||
pip install colossalai==0.3.7
|
||||
```
|
||||
|
||||
Install OpenDiT:
|
||||
|
||||
```shell
|
||||
git clone https://github.com/NUS-HPC-AI-Lab/OpenDiT
|
||||
cd OpenDiT
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
|
||||
## Usage
|
||||
|
||||
OpenDiT fully supports the following models, including training and inference, which align with the original methods. Through our novel techniques, we enable these models to run faster and consume less memory. Here's how you can use them:
|
||||
|
||||
| Model | Train | Inference | Optimize | Usage |
|
||||
| ------ | :------: | :------: | :------: | :------: |
|
||||
| DiT[[source](https://github.com/facebookresearch/DiT)]| ✅ | ✅ | ✅ | [Doc](./docs/dit.md)
|
||||
| Open-Sora[[source](https://github.com/hpcaitech/Open-Sora)]| 🟡 | ✅ | ✅ | [Doc](./docs/opensora.md)
|
||||
| Latte[[source](https://github.com/Vchitect/Latte)]| ❌ | ✅ | ✅ | [Doc](./docs/latte.md)
|
||||
| Open-Sora-Plan[[source](https://github.com/PKU-YuanGroup/Open-Sora-Plan)]| ❌ | ✅ | ✅ | [Doc](./docs/opensora_plan.md)
|
||||
|
||||
## Technique Overview
|
||||
|
||||
### Pyramid Attention Broadcast (PAB) [[blog](https://arxiv.org/abs/2403.10266)][[doc](./docs/pab.md)]
|
||||
|
||||
Real-Time Video Generation with Pyramid Attention Broadcast
|
||||
|
||||
Authors: [Xuanlei Zhao](https://oahzxl.github.io/)<sup>1*</sup>, [Xiaolong Jin]()<sup>2*</sup>, [Kai Wang](https://kaiwang960112.github.io/)<sup>1*</sup>, and [Yang You](https://www.comp.nus.edu.sg/~youy/)<sup>1</sup> (* indicates equal contribution)
|
||||
|
||||
<sup>1</sup>National University of Singapore, <sup>2</sup>Purdue University
|
||||
|
||||

|
||||
|
||||
PAB is the first approach to achieve <b>real-time</b> DiT-based video generation, delivering <b>lossless quality</b> without <b>requiring any training</b>.
|
||||
|
||||
By mitigating redundant attention computation, PAB achieves up to 21.6 FPS with 10.6x acceleration, without sacrificing quality across popular DiT-based video generation models including Open-Sora, Open-Sora-Plan, and Latte.
|
||||
|
||||
Notably, as a training-free approach, PAB can enpower any future DiT-based video generation models with real-time capabilities.
|
||||
|
||||
See its detail and usage [here](./docs/pab.md).
|
||||
|
||||
----
|
||||
|
||||
### Dyanmic Sequence Parallelism (DSP) [[paper](https://arxiv.org/abs/2403.10266)][[doc](./docs/dsp.md)]
|
||||
|
||||

|
||||
|
||||
DSP is a novel, elegant and super efficient sequence parallelism for [OpenSora](https://github.com/hpcaitech/Open-Sora), [Latte](https://github.com/Vchitect/Latte) and other multi-dimensional transformer architecture.
|
||||
|
||||
It achieves **3x** speed for training and **2x** speed for inference in OpenSora compared with sota sequence parallelism ([DeepSpeed Ulysses](https://arxiv.org/abs/2309.14509)). For a 10s (80 frames) of 512x512 video, the inference latency of OpenSora is:
|
||||
|
||||
| Method | 1xH800 | 8xH800 (DS Ulysses) | 8xH800 (DSP) |
|
||||
| ------ | ------ | ------ | ------ |
|
||||
| Latency(s) | 106 | 45 | 22 |
|
||||
|
||||
See its detail and usage [here](./docs/dsp.md).
|
||||
|
||||
----
|
||||
|
||||
## DiT Reproduction Result
|
||||
|
||||
We have trained DiT using the origin method with OpenDiT to verify our accuracy. We have trained the model from scratch on ImageNet for 80k steps on 8xA100. Here are some results generated by our trained DiT:
|
||||
|
||||

|
||||
|
||||
Our loss also aligns with the results listed in the paper:
|
||||
|
||||

|
||||
|
||||
To reproduce our results, you can follow our [instruction](./docs/dit.md/#reproduction
|
||||
).
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
Thanks [Xuanlei Zhao](https://oahzxl.github.io/), [Zhongkai Zhao](https://www.linkedin.com/in/zhongkai-zhao-kk2000/), [Ziming Liu](https://maruyamaaya.github.io/), [Haotian Zhou](https://github.com/ht-zhou), [Qianli Ma](https://fazzie-key.cool/about/index.html), [Yang You](https://www.comp.nus.edu.sg/~youy/), [Xiaolong Jin](), [Kai Wang](https://kaiwang960112.github.io/) for their contributions. We also extend our gratitude to [Zangwei Zheng](https://zhengzangw.github.io/), [Shenggan Cheng](https://shenggan.github.io/), [Fuzhao Xue](https://xuefuzhao.github.io/), [Shizun Wang](https://littlepure2333.github.io/home/), [Yuchao Gu](https://ycgu.site/), [Shenggui Li](https://franklee.xyz/), and [Haofan Wang](https://haofanwang.github.io/) for their invaluable advice.
|
||||
|
||||
This codebase borrows from:
|
||||
* [Open-Sora](https://github.com/hpcaitech/Open-Sora): Democratizing Efficient Video Production for All.
|
||||
* [DiT](https://github.com/facebookresearch/DiT): Scalable Diffusion Models with Transformers.
|
||||
* [PixArt](https://github.com/PixArt-alpha/PixArt-alpha): An open-source DiT-based text-to-image model.
|
||||
* [Latte](https://github.com/Vchitect/Latte): An attempt to efficiently train DiT for video.
|
||||
|
||||
## Contributing
|
||||
|
||||
If you encounter problems using OpenDiT or have a feature request, feel free to create an issue! We also welcome pull requests from the community.
|
||||
|
||||
## Citation
|
||||
|
||||
```
|
||||
@misc{zhao2024opendit,
|
||||
author = {Xuanlei Zhao, Zhongkai Zhao, Ziming Liu, Haotian Zhou, Qianli Ma, and Yang You},
|
||||
title = {OpenDiT: An Easy, Fast and Memory-Efficient System for DiT Training and Inference},
|
||||
year = {2024},
|
||||
publisher = {GitHub},
|
||||
journal = {GitHub repository},
|
||||
howpublished = {\url{https://github.com/NUS-HPC-AI-Lab/OpenDiT}},
|
||||
}
|
||||
|
||||
@misc{zhao2024dsp,
|
||||
title={DSP: Dynamic Sequence Parallelism for Multi-Dimensional Transformers},
|
||||
author={Xuanlei Zhao and Shenggan Cheng and Zangwei Zheng and Zheming Yang and Ziming Liu and Yang You},
|
||||
year={2024},
|
||||
eprint={2403.10266},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.DC}
|
||||
}
|
||||
```
|
||||
|
||||
## Star History
|
||||
|
||||
[](https://star-history.com/#NUS-HPC-AI-Lab/OpenDiT&Date)
|
||||
Executable
+3
@@ -0,0 +1,3 @@
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
Executable
+41
@@ -0,0 +1,41 @@
|
||||
# path:
|
||||
save_img_path: "./samples/latte/"
|
||||
pretrained_model_path: "maxin-cn/Latte-1"
|
||||
|
||||
# model config:
|
||||
model: LatteT2V
|
||||
video_length: 16
|
||||
image_size: [512, 512]
|
||||
# # beta schedule
|
||||
beta_start: 0.0001
|
||||
beta_end: 0.02
|
||||
beta_schedule: "linear"
|
||||
variance_type: "learned_range"
|
||||
|
||||
# model speedup
|
||||
use_compile: False
|
||||
use_fp16: True
|
||||
|
||||
# sample config:
|
||||
seed: 0
|
||||
run_time: 0
|
||||
guidance_scale: 7.5
|
||||
sample_method: 'DDIM'
|
||||
num_sampling_steps: 50
|
||||
enable_temporal_attentions: True
|
||||
enable_vae_temporal_decoder: True # use temporal vae decoder from SVD, maybe reduce the video flicker (It's not widely tested)
|
||||
|
||||
text_prompt: [
|
||||
"Time Lapse of the rising sun over a tree in an open rural landscape, with clouds in the blue sky beautifully playing with the rays of light",
|
||||
"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. The sun is shining brightly, casting a warm glow on the flowers and highlighting their intricate details. The video is shot from a low angle, looking up at the sunflowers, which adds a sense of grandeur and awe to the scene. The sunflowers are the main focus of the video, with no other objects or people present. The video is a celebration of nature's beauty and the simple joy of a sunny day in the countryside.",
|
||||
"Snow falling over multiple houses and trees on winter landscape against night sky. christmas festivity and celebration concept.",
|
||||
"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.",
|
||||
"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.",
|
||||
'An epic tornado attacking above aglowing city at night.',
|
||||
'Slow pan upward of blazing oak fire in an indoor fireplace.',
|
||||
'Sunset over the sea.',
|
||||
'A dog in astronaut suit and sunglasses floating in space.',
|
||||
"The niagara river in new york has very rough white water that you aren't allowed raft or kayak",
|
||||
"Peru - july 03, 2014: big wave world tour- wave rolling in off the coast of peru.",
|
||||
"Waterfall names 'namtok chokkadin' at thong pha phum national park, kanchanaburi province, thailand, panning tilt up in low angle view",
|
||||
]
|
||||
Executable
+54
@@ -0,0 +1,54 @@
|
||||
# path:
|
||||
save_img_path: "./samples/latte/"
|
||||
pretrained_model_path: "maxin-cn/Latte-1"
|
||||
|
||||
# model config:
|
||||
model: LatteT2V
|
||||
video_length: 16
|
||||
image_size: [512, 512]
|
||||
# # beta schedule
|
||||
beta_start: 0.0001
|
||||
beta_end: 0.02
|
||||
beta_schedule: "linear"
|
||||
variance_type: "learned_range"
|
||||
|
||||
# model speedup
|
||||
use_compile: False
|
||||
use_fp16: True
|
||||
|
||||
# sample config:
|
||||
seed: 0
|
||||
run_time: 0
|
||||
guidance_scale: 7.5
|
||||
sample_method: 'DDIM'
|
||||
num_sampling_steps: 50
|
||||
enable_temporal_attentions: True
|
||||
enable_vae_temporal_decoder: True # use temporal vae decoder from SVD, maybe reduce the video flicker (It's not widely tested)
|
||||
|
||||
text_prompt: [
|
||||
"Time Lapse of the rising sun over a tree in an open rural landscape, with clouds in the blue sky beautifully playing with the rays of light",
|
||||
"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. The sun is shining brightly, casting a warm glow on the flowers and highlighting their intricate details. The video is shot from a low angle, looking up at the sunflowers, which adds a sense of grandeur and awe to the scene. The sunflowers are the main focus of the video, with no other objects or people present. The video is a celebration of nature's beauty and the simple joy of a sunny day in the countryside.",
|
||||
"Snow falling over multiple houses and trees on winter landscape against night sky. christmas festivity and celebration concept.",
|
||||
"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.",
|
||||
"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.",
|
||||
'An epic tornado attacking above aglowing city at night.',
|
||||
'Slow pan upward of blazing oak fire in an indoor fireplace.',
|
||||
'Sunset over the sea.',
|
||||
'A dog in astronaut suit and sunglasses floating in space.',
|
||||
"The niagara river in new york has very rough white water that you aren't allowed raft or kayak",
|
||||
"Peru - july 03, 2014: big wave world tour- wave rolling in off the coast of peru.",
|
||||
"Waterfall names 'namtok chokkadin' at thong pha phum national park, kanchanaburi province, thailand, panning tilt up in low angle view",
|
||||
]
|
||||
|
||||
# pab
|
||||
spatial_broadcast: True
|
||||
spatial_threshold: [100, 800]
|
||||
spatial_gap: 2
|
||||
temporal_broadcast: True
|
||||
temporal_threshold: [100, 800]
|
||||
temporal_gap: 4
|
||||
cross_broadcast: True
|
||||
cross_threshold: [80, 900]
|
||||
cross_gap: 7
|
||||
# diffusion_skip: True
|
||||
# diffusion_skip_timestep: [1,1,1,0,0,0,0,0,0,0]
|
||||
Executable
+34
@@ -0,0 +1,34 @@
|
||||
resolution: "480p"
|
||||
aspect_ratio: "9:16"
|
||||
num_frames: 48
|
||||
fps: 24
|
||||
frame_interval: 1
|
||||
|
||||
seed: 42
|
||||
multi_resolution: "STDiT2"
|
||||
dtype: "bf16"
|
||||
condition_frame_length: 5
|
||||
align: 5
|
||||
num_sampling_steps: 30
|
||||
cfg_scale: 7.0
|
||||
aes: 6.5
|
||||
|
||||
prompt_as_path: true
|
||||
prompt: [
|
||||
"Time Lapse of the rising sun over a tree in an open rural landscape, with clouds in the blue sky beautifully playing with the rays of light",
|
||||
"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. The sun is shining brightly, casting a warm glow on the flowers and highlighting their intricate details. The video is shot from a low angle, looking up at the sunflowers, which adds a sense of grandeur and awe to the scene. The sunflowers are the main focus of the video, with no other objects or people present. The video is a celebration of nature's beauty and the simple joy of a sunny day in the countryside.",
|
||||
"Snow falling over multiple houses and trees on winter landscape against night sky. christmas festivity and celebration concept.",
|
||||
"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.",
|
||||
"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.",
|
||||
'An epic tornado attacking above aglowing city at night.',
|
||||
'Slow pan upward of blazing oak fire in an indoor fireplace.',
|
||||
'Sunset over the sea.',
|
||||
'A dog in astronaut suit and sunglasses floating in space.',
|
||||
"The niagara river in new york has very rough white water that you aren't allowed raft or kayak",
|
||||
"Peru - july 03, 2014: big wave world tour- wave rolling in off the coast of peru.",
|
||||
"Waterfall names 'namtok chokkadin' at thong pha phum national park, kanchanaburi province, thailand, panning tilt up in low angle view",
|
||||
]
|
||||
|
||||
# speedup
|
||||
flash_attn: False # turn on for faster inference
|
||||
enable_t5_speedup: False # turn on for less memory usage
|
||||
Executable
+47
@@ -0,0 +1,47 @@
|
||||
resolution: "480p"
|
||||
aspect_ratio: "9:16"
|
||||
num_frames: 48
|
||||
fps: 24
|
||||
frame_interval: 1
|
||||
|
||||
seed: 42
|
||||
multi_resolution: "STDiT2"
|
||||
dtype: "bf16"
|
||||
condition_frame_length: 5
|
||||
align: 5
|
||||
num_sampling_steps: 30
|
||||
cfg_scale: 7.0
|
||||
aes: 6.5
|
||||
|
||||
prompt_as_path: true
|
||||
prompt: [
|
||||
"Time Lapse of the rising sun over a tree in an open rural landscape, with clouds in the blue sky beautifully playing with the rays of light",
|
||||
"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. The sun is shining brightly, casting a warm glow on the flowers and highlighting their intricate details. The video is shot from a low angle, looking up at the sunflowers, which adds a sense of grandeur and awe to the scene. The sunflowers are the main focus of the video, with no other objects or people present. The video is a celebration of nature's beauty and the simple joy of a sunny day in the countryside.",
|
||||
"Snow falling over multiple houses and trees on winter landscape against night sky. christmas festivity and celebration concept.",
|
||||
"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.",
|
||||
"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.",
|
||||
'An epic tornado attacking above aglowing city at night.',
|
||||
'Slow pan upward of blazing oak fire in an indoor fireplace.',
|
||||
'Sunset over the sea.',
|
||||
'A dog in astronaut suit and sunglasses floating in space.',
|
||||
"The niagara river in new york has very rough white water that you aren't allowed raft or kayak",
|
||||
"Peru - july 03, 2014: big wave world tour- wave rolling in off the coast of peru.",
|
||||
"Waterfall names 'namtok chokkadin' at thong pha phum national park, kanchanaburi province, thailand, panning tilt up in low angle view",
|
||||
]
|
||||
|
||||
# speedup
|
||||
flash_attn: True
|
||||
enable_t5_speedup: True
|
||||
|
||||
# pab
|
||||
spatial_broadcast: True
|
||||
spatial_threshold: [540, 940]
|
||||
spatial_gap: 2
|
||||
temporal_broadcast: True
|
||||
temporal_threshold: [540, 940]
|
||||
temporal_gap: 4
|
||||
cross_broadcast: True
|
||||
cross_threshold: [540, 940]
|
||||
cross_gap: 6
|
||||
# diffusion_skip: True
|
||||
# diffusion_skip_timestep: [1,1,1,0,0,0,0,0,0,0]
|
||||
Executable
+27
@@ -0,0 +1,27 @@
|
||||
model_path: LanguageBind/Open-Sora-Plan-v1.1.0
|
||||
version: 65x512x512
|
||||
num_frames: 65
|
||||
height: 512
|
||||
width: 512
|
||||
cache_dir: "./cache_dir"
|
||||
text_encoder_name: DeepFloyd/t5-v1_1-xxl
|
||||
text_prompt: [
|
||||
"Time Lapse of the rising sun over a tree in an open rural landscape, with clouds in the blue sky beautifully playing with the rays of light",
|
||||
"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. The sun is shining brightly, casting a warm glow on the flowers and highlighting their intricate details. The video is shot from a low angle, looking up at the sunflowers, which adds a sense of grandeur and awe to the scene. The sunflowers are the main focus of the video, with no other objects or people present. The video is a celebration of nature's beauty and the simple joy of a sunny day in the countryside.",
|
||||
"Snow falling over multiple houses and trees on winter landscape against night sky. christmas festivity and celebration concept.",
|
||||
"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.",
|
||||
"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.",
|
||||
'An epic tornado attacking above aglowing city at night.',
|
||||
'Slow pan upward of blazing oak fire in an indoor fireplace.',
|
||||
'Sunset over the sea.',
|
||||
'A dog in astronaut suit and sunglasses floating in space.',
|
||||
"The niagara river in new york has very rough white water that you aren't allowed raft or kayak",
|
||||
"Peru - july 03, 2014: big wave world tour- wave rolling in off the coast of peru.",
|
||||
"Waterfall names 'namtok chokkadin' at thong pha phum national park, kanchanaburi province, thailand, panning tilt up in low angle view",
|
||||
]
|
||||
ae: CausalVAEModel_4x8x8
|
||||
save_img_path: "./samples/opensora_plan"
|
||||
fps: 24
|
||||
guidance_scale: 7.5
|
||||
num_sampling_steps: 150
|
||||
enable_tiling: True
|
||||
Executable
+40
@@ -0,0 +1,40 @@
|
||||
model_path: LanguageBind/Open-Sora-Plan-v1.1.0
|
||||
version: 65x512x512
|
||||
num_frames: 65
|
||||
height: 512
|
||||
width: 512
|
||||
cache_dir: "./cache_dir"
|
||||
text_encoder_name: DeepFloyd/t5-v1_1-xxl
|
||||
text_prompt: [
|
||||
"Time Lapse of the rising sun over a tree in an open rural landscape, with clouds in the blue sky beautifully playing with the rays of light",
|
||||
"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. The sun is shining brightly, casting a warm glow on the flowers and highlighting their intricate details. The video is shot from a low angle, looking up at the sunflowers, which adds a sense of grandeur and awe to the scene. The sunflowers are the main focus of the video, with no other objects or people present. The video is a celebration of nature's beauty and the simple joy of a sunny day in the countryside.",
|
||||
"Snow falling over multiple houses and trees on winter landscape against night sky. christmas festivity and celebration concept.",
|
||||
"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.",
|
||||
"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.",
|
||||
'An epic tornado attacking above aglowing city at night.',
|
||||
'Slow pan upward of blazing oak fire in an indoor fireplace.',
|
||||
'Sunset over the sea.',
|
||||
'A dog in astronaut suit and sunglasses floating in space.',
|
||||
"The niagara river in new york has very rough white water that you aren't allowed raft or kayak",
|
||||
"Peru - july 03, 2014: big wave world tour- wave rolling in off the coast of peru.",
|
||||
"Waterfall names 'namtok chokkadin' at thong pha phum national park, kanchanaburi province, thailand, panning tilt up in low angle view",
|
||||
]
|
||||
ae: CausalVAEModel_4x8x8
|
||||
save_img_path: "./samples/opensora_plan"
|
||||
fps: 24
|
||||
guidance_scale: 7.5
|
||||
num_sampling_steps: 150
|
||||
enable_tiling: True
|
||||
|
||||
# pab
|
||||
spatial_broadcast: True
|
||||
spatial_threshold: [100, 800]
|
||||
spatial_gap: 2
|
||||
temporal_broadcast: True
|
||||
temporal_threshold: [100, 800]
|
||||
temporal_gap: 4
|
||||
cross_broadcast: True
|
||||
cross_threshold: [100, 850]
|
||||
cross_gap: 6
|
||||
# diffusion_skip: True
|
||||
# diffusion_skip_timestep: [3,3,3,0,0,0,0,0,0,0]
|
||||
Executable
+32
@@ -0,0 +1,32 @@
|
||||
{
|
||||
"_name_or_path": "t5-v1_1-xxl-encoder-bf16",
|
||||
"architectures": [
|
||||
"T5EncoderModel"
|
||||
],
|
||||
"classifier_dropout": 0.0,
|
||||
"d_ff": 10240,
|
||||
"d_kv": 64,
|
||||
"d_model": 4096,
|
||||
"decoder_start_token_id": 0,
|
||||
"dense_act_fn": "gelu_new",
|
||||
"dropout_rate": 0.1,
|
||||
"eos_token_id": 1,
|
||||
"feed_forward_proj": "gated-gelu",
|
||||
"initializer_factor": 1.0,
|
||||
"is_encoder_decoder": true,
|
||||
"is_gated_act": true,
|
||||
"layer_norm_epsilon": 1e-06,
|
||||
"model_type": "t5",
|
||||
"num_decoder_layers": 24,
|
||||
"num_heads": 64,
|
||||
"num_layers": 24,
|
||||
"output_past": true,
|
||||
"pad_token_id": 0,
|
||||
"relative_attention_max_distance": 128,
|
||||
"relative_attention_num_buckets": 32,
|
||||
"tie_word_embeddings": false,
|
||||
"torch_dtype": "bfloat16",
|
||||
"transformers_version": "4.40.1",
|
||||
"use_cache": true,
|
||||
"vocab_size": 32128
|
||||
}
|
||||
@@ -0,0 +1,329 @@
|
||||
{
|
||||
"last_node_id": 15,
|
||||
"last_link_id": 18,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 3,
|
||||
"type": "VHS_VideoCombine",
|
||||
"pos": [
|
||||
580,
|
||||
-350
|
||||
],
|
||||
"size": [
|
||||
669.1505737304688,
|
||||
671.4437209795105
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 18
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "VHS_AUDIO",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "Filenames",
|
||||
"type": "VHS_FILENAMES",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_VideoCombine"
|
||||
},
|
||||
"widgets_values": {
|
||||
"frame_rate": 16,
|
||||
"loop_count": 0,
|
||||
"filename_prefix": "Open-Sora-PAB",
|
||||
"format": "video/h264-mp4",
|
||||
"pix_fmt": "yuv420p",
|
||||
"crf": 19,
|
||||
"save_metadata": true,
|
||||
"pingpong": false,
|
||||
"save_output": false,
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "Open-Sora-PAB_00003.mp4",
|
||||
"subfolder": "",
|
||||
"type": "temp",
|
||||
"format": "video/h264-mp4"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 11,
|
||||
"type": "OpenDiTConditioning",
|
||||
"pos": [
|
||||
-346,
|
||||
-181
|
||||
],
|
||||
"size": {
|
||||
"0": 470,
|
||||
"1": 308.20001220703125
|
||||
},
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "opendit_t5_encoder",
|
||||
"type": "OPENDITT5",
|
||||
"link": 14,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "opendit_cond",
|
||||
"type": "OPENDITCOND",
|
||||
"links": [
|
||||
17
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "OpenDiTConditioning"
|
||||
},
|
||||
"widgets_values": [
|
||||
"video of a waterfall",
|
||||
"",
|
||||
6.5,
|
||||
0,
|
||||
"bf16"
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 10,
|
||||
"type": "DownloadAndLoadOpenSoraModel",
|
||||
"pos": [
|
||||
-689,
|
||||
-465
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 82
|
||||
},
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "opendit_model",
|
||||
"type": "OPENDITMODEL",
|
||||
"links": [
|
||||
15
|
||||
],
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "DownloadAndLoadOpenSoraModel"
|
||||
},
|
||||
"widgets_values": [
|
||||
"hpcai-tech/OpenSora-STDiT-v3",
|
||||
"bf16"
|
||||
],
|
||||
"color": "#323",
|
||||
"bgcolor": "#535"
|
||||
},
|
||||
{
|
||||
"id": 15,
|
||||
"type": "OpenDiTSampler",
|
||||
"pos": [
|
||||
168,
|
||||
-294
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 290
|
||||
},
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "opendit_model",
|
||||
"type": "OPENDITMODEL",
|
||||
"link": 15
|
||||
},
|
||||
{
|
||||
"name": "opendit_vae",
|
||||
"type": "OPENDITVAE",
|
||||
"link": 16
|
||||
},
|
||||
{
|
||||
"name": "opendit_cond",
|
||||
"type": "OPENDITCOND",
|
||||
"link": 17
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
18
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "OpenDiTSampler"
|
||||
},
|
||||
"widgets_values": [
|
||||
25,
|
||||
426,
|
||||
240,
|
||||
863675962091352,
|
||||
"randomize",
|
||||
25,
|
||||
7,
|
||||
24,
|
||||
false
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 13,
|
||||
"type": "DownloadAndLoadOpenSoraVAE",
|
||||
"pos": [
|
||||
-690,
|
||||
-330
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 82
|
||||
},
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "opendit_vae",
|
||||
"type": "OPENDITVAE",
|
||||
"links": [
|
||||
16
|
||||
],
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "DownloadAndLoadOpenSoraVAE"
|
||||
},
|
||||
"widgets_values": [
|
||||
"hpcai-tech/OpenSora-VAE-v1.2",
|
||||
"bf16"
|
||||
],
|
||||
"color": "#322",
|
||||
"bgcolor": "#533"
|
||||
},
|
||||
{
|
||||
"id": 14,
|
||||
"type": "DownloadAndLoadOpenDiTT5Model",
|
||||
"pos": [
|
||||
-696,
|
||||
-186
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 82
|
||||
},
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "opendit_t5_encoder",
|
||||
"type": "OPENDITT5",
|
||||
"links": [
|
||||
14
|
||||
],
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "DownloadAndLoadOpenDiTT5Model"
|
||||
},
|
||||
"widgets_values": [
|
||||
"city96/t5-v1_1-xxl-encoder-bf16",
|
||||
"bf16"
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
14,
|
||||
14,
|
||||
0,
|
||||
11,
|
||||
0,
|
||||
"OPENDITT5"
|
||||
],
|
||||
[
|
||||
15,
|
||||
10,
|
||||
0,
|
||||
15,
|
||||
0,
|
||||
"OPENDITMODEL"
|
||||
],
|
||||
[
|
||||
16,
|
||||
13,
|
||||
0,
|
||||
15,
|
||||
1,
|
||||
"OPENDITVAE"
|
||||
],
|
||||
[
|
||||
17,
|
||||
11,
|
||||
0,
|
||||
15,
|
||||
2,
|
||||
"OPENDITCOND"
|
||||
],
|
||||
[
|
||||
18,
|
||||
15,
|
||||
0,
|
||||
3,
|
||||
0,
|
||||
"IMAGE"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.683013455365071,
|
||||
"offset": {
|
||||
"0": 781.8367309570312,
|
||||
"1": 570.311767578125
|
||||
}
|
||||
}
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -0,0 +1,364 @@
|
||||
import torch
|
||||
import os
|
||||
import sys
|
||||
import gc
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
sys.path.append(script_directory)
|
||||
|
||||
from opendit.core.pab_mgr import set_pab_manager
|
||||
#from opendit.core.parallel_mgr import enable_sequence_parallel, set_parallel_manager
|
||||
from opendit.models.opensora import RFLOW, OpenSoraVAE_V1_2, STDiT3_XL_2, T5Encoder, text_preprocessing
|
||||
from opendit.models.opensora.inference_utils import (
|
||||
append_score_to_prompts,
|
||||
extract_prompts_loop,
|
||||
merge_prompt,
|
||||
prepare_multi_resolution_info,
|
||||
split_prompt,
|
||||
)
|
||||
|
||||
import comfy.model_management as mm
|
||||
from comfy.utils import ProgressBar, load_torch_file
|
||||
import folder_paths
|
||||
|
||||
try:
|
||||
from flash_attn import flash_attn_varlen_func
|
||||
FLASH_ATTN_AVAILABLE = True
|
||||
print("Flash Attention is available")
|
||||
except:
|
||||
FLASH_ATTN_AVAILABLE = False
|
||||
print("WARNING! Flash Attention is not available, using much slower torch SDP attention")
|
||||
|
||||
class DownloadAndLoadOpenSoraModel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"model": (
|
||||
[
|
||||
'hpcai-tech/OpenSora-STDiT-v3'
|
||||
],
|
||||
),
|
||||
"precision": (['fp16','bf16','fp32'],
|
||||
{
|
||||
"default": 'bf16'
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("OPENDITMODEL",)
|
||||
RETURN_NAMES = ("opendit_model",)
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "OpenDitWrapper"
|
||||
|
||||
def loadmodel(self, model, precision):
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
|
||||
model_name = model.rsplit('/', 1)[-1]
|
||||
model_path = os.path.join(folder_paths.models_dir, "opensora", model_name)
|
||||
|
||||
if not os.path.exists(model_path):
|
||||
print(f"Downloading OpenSora model to: {model_path}")
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download(repo_id=model,
|
||||
ignore_patterns=['*ema*'],
|
||||
local_dir=model_path,
|
||||
local_dir_use_symlinks=False)
|
||||
|
||||
if not hasattr(self, "model"):
|
||||
print("Loading STDiT...")
|
||||
self.model = (
|
||||
STDiT3_XL_2(
|
||||
from_pretrained=model_path,
|
||||
qk_norm=True,
|
||||
enable_flash_attn=FLASH_ATTN_AVAILABLE,
|
||||
enable_layernorm_kernel=True,
|
||||
#input_size=latent_size,
|
||||
in_channels=4,
|
||||
caption_channels=4096,
|
||||
model_max_length=300
|
||||
).to(offload_device, dtype).eval()
|
||||
)
|
||||
|
||||
mm.soft_empty_cache()
|
||||
|
||||
opendit_model = {
|
||||
'model': self.model,
|
||||
'dtype': dtype
|
||||
}
|
||||
|
||||
return (opendit_model,)
|
||||
|
||||
class DownloadAndLoadOpenSoraVAE:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"model": (
|
||||
[
|
||||
'hpcai-tech/OpenSora-VAE-v1.2'
|
||||
],
|
||||
),
|
||||
"precision": (['fp16','bf16','fp32'],
|
||||
{
|
||||
"default": 'bf16'
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("OPENDITVAE",)
|
||||
RETURN_NAMES = ("opendit_vae",)
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "OpenDitWrapper"
|
||||
|
||||
def loadmodel(self, model, precision):
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
|
||||
model_name = model.rsplit('/', 1)[-1]
|
||||
model_path = os.path.join(folder_paths.models_dir, "opensora", model_name)
|
||||
|
||||
if not os.path.exists(model_path):
|
||||
print(f"Downloading OpenSora model to: {model_path}")
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download(repo_id=model,
|
||||
ignore_patterns=['*ema*'],
|
||||
local_dir=model_path,
|
||||
local_dir_use_symlinks=False)
|
||||
|
||||
if not hasattr(self, "vae"):
|
||||
print("Loading VAE...")
|
||||
self.vae = (
|
||||
OpenSoraVAE_V1_2(
|
||||
from_pretrained="hpcai-tech/OpenSora-VAE-v1.2",
|
||||
micro_frame_size=17,
|
||||
micro_batch_size=4,
|
||||
).to(offload_device, dtype).eval()
|
||||
)
|
||||
|
||||
mm.soft_empty_cache()
|
||||
|
||||
opendit_model = {
|
||||
'model': self.vae,
|
||||
'dtype': dtype
|
||||
}
|
||||
|
||||
return (opendit_model,)
|
||||
|
||||
class DownloadAndLoadOpenDiTT5Model:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"model": (
|
||||
[
|
||||
'city96/t5-v1_1-xxl-encoder-bf16'
|
||||
],
|
||||
),
|
||||
"precision": (['fp16','bf16','fp32'],
|
||||
{
|
||||
"default": 'bf16'
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("OPENDITT5",)
|
||||
RETURN_NAMES = ("opendit_t5_encoder",)
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "OpenDitWrapper"
|
||||
|
||||
def loadmodel(self, model, precision):
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
|
||||
model_name = model.rsplit('/', 1)[-1]
|
||||
model_path = os.path.join(folder_paths.models_dir, "t5", model_name)
|
||||
|
||||
if not os.path.exists(model_path):
|
||||
print(f"Downloading OpenSora model to: {model_path}")
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download(repo_id=model,
|
||||
ignore_patterns=['*ema*'],
|
||||
local_dir=model_path,
|
||||
local_dir_use_symlinks=False)
|
||||
|
||||
|
||||
if not hasattr(self, "text_encoder"):
|
||||
print("Loading Text Encoder...")
|
||||
self.text_encoder = T5Encoder(
|
||||
from_pretrained=model_path, model_max_length=300, device=device, dtype=dtype, shardformer=False
|
||||
)
|
||||
|
||||
mm.soft_empty_cache()
|
||||
|
||||
return (self.text_encoder,)
|
||||
|
||||
class OpenDiTConditioning:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"opendit_t5_encoder": ("OPENDITT5",),
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"camera_prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"aesthetic_score": ("FLOAT", {"default": 6.5, "min": 0.0, "max": 100.0, "step": 0.1}),
|
||||
"flow_score": ("FLOAT", {"default": 0.0}, {"min": 0.0, "max": 100.0, "step": 0.1}),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("OPENDITCOND",)
|
||||
RETURN_NAMES =("opendit_cond",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "OpenDiTWrapper"
|
||||
|
||||
def process(self, opendit_t5_encoder, prompt, camera_prompt, aesthetic_score, flow_score, keep_model_loaded=False):
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
self.text_encoder = opendit_t5_encoder
|
||||
|
||||
print("process prompt step by step...")
|
||||
# == process prompt step by step ==
|
||||
# 0. split prompt
|
||||
prompt_segment_list, loop_idx_list = split_prompt(prompt)
|
||||
|
||||
# 1. append score
|
||||
prompt_segment_list = append_score_to_prompts(
|
||||
prompt_segment_list,
|
||||
aes=aesthetic_score if aesthetic_score > 0 else None,
|
||||
flow=flow_score if flow_score > 0 else None,
|
||||
camera_motion=camera_prompt if camera_prompt != "" else None,
|
||||
)
|
||||
|
||||
# 2. clean prompt with T5
|
||||
prompt_segment_list = [text_preprocessing(prompt) for prompt in prompt_segment_list]
|
||||
|
||||
# 3. merge to obtain the final prompt
|
||||
final_prompt = merge_prompt(prompt_segment_list, loop_idx_list)
|
||||
final_prompt_loop = extract_prompts_loop([final_prompt], 0)
|
||||
print("final_prompt_loop: ", final_prompt_loop)
|
||||
|
||||
self.text_encoder.t5.model.to(device)
|
||||
encoded_prompt = self.text_encoder.encode(final_prompt_loop)
|
||||
if not keep_model_loaded:
|
||||
self.text_encoder.t5.model.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
return (encoded_prompt,)
|
||||
class OpenDiTSampler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"opendit_model": ("OPENDITMODEL",),
|
||||
"opendit_vae": ("OPENDITVAE",),
|
||||
"opendit_cond": ("OPENDITCOND",),
|
||||
"num_frames": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}),
|
||||
"width": ("INT", {"default": 426, "min": 1, "max": 2048, "step": 1}),
|
||||
"height": ("INT", {"default": 240, "min": 1, "max": 2048, "step": 1}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"steps": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}),
|
||||
"cfg": ("FLOAT", {"default": 7.0, "min": 0.0, "max": 20.0, "step": 0.01}),
|
||||
"fps": ("INT", {"default": 24, "min": 1, "max": 60, "step": 1}),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES =("images",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "OpenDiTWrapper"
|
||||
|
||||
def process(self, opendit_model, opendit_vae, opendit_cond, num_frames, width, height, seed, steps, cfg, fps, keep_model_loaded=False):
|
||||
device = mm.get_torch_device()
|
||||
dtype = opendit_model['dtype']
|
||||
offload_device = mm.unet_offload_device()
|
||||
self.model = opendit_model['model']
|
||||
self.vae = opendit_vae['model']
|
||||
|
||||
set_pab_manager(
|
||||
steps=steps,
|
||||
cross_broadcast=True,
|
||||
cross_threshold=[540, 940],
|
||||
cross_gap=6,
|
||||
spatial_broadcast=True,
|
||||
spatial_threshold=[540, 940],
|
||||
spatial_gap=2,
|
||||
temporal_broadcast=True,
|
||||
temporal_threshold=[540, 940],
|
||||
temporal_gap=4,
|
||||
#diffusion_skip=6
|
||||
#diffusion_skip_timestep= [1,1,1,0,0,0,0,0,0,0]
|
||||
)
|
||||
|
||||
image_size = (height, width)
|
||||
input_size = (num_frames, *image_size)
|
||||
latent_size = self.vae.get_latent_size(input_size)
|
||||
|
||||
scheduler = RFLOW(use_timestep_transform=True, num_sampling_steps=steps, cfg_scale=cfg)
|
||||
|
||||
print("Sampling...")
|
||||
# == sampling ==
|
||||
torch.manual_seed(seed)
|
||||
z = torch.randn(1, self.vae.out_channels, *latent_size, device=device, dtype=dtype)
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
multi_resolution = "STDiT2"
|
||||
additional_args = prepare_multi_resolution_info(
|
||||
multi_resolution, 1, image_size, num_frames, fps, device, dtype
|
||||
)
|
||||
print("additional_args: ", additional_args)
|
||||
final_cond = opendit_cond.copy()
|
||||
final_cond.update(additional_args)
|
||||
|
||||
self.model.to(device)
|
||||
|
||||
y_null = self.model.y_embedder.y_embedding[None].repeat(1, 1, 1)[:, None]
|
||||
final_cond["y"] = torch.cat([final_cond["y"], y_null], 0)
|
||||
|
||||
samples = scheduler.sample(
|
||||
self.model,
|
||||
final_cond,
|
||||
z=z,
|
||||
device=device,
|
||||
progress=True,
|
||||
additional_args=additional_args,
|
||||
# mask=masks, # Adjust or omit based on your needs
|
||||
)
|
||||
if not keep_model_loaded:
|
||||
self.model.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
self.vae.to(device)
|
||||
samples = self.vae.decode(samples.to(dtype), num_frames=num_frames)
|
||||
self.vae.to(offload_device)
|
||||
|
||||
samples = samples.squeeze(0).permute(1, 2, 3, 0).float().cpu()
|
||||
normalized_tensor = torch.clamp(samples, -1, 1)
|
||||
|
||||
tensor_min = normalized_tensor.min()
|
||||
tensor_max = normalized_tensor.max()
|
||||
normalized_tensor = (samples - tensor_min) / (tensor_max - tensor_min)
|
||||
normalized_tensor = torch.clamp(normalized_tensor, 0, 1)
|
||||
|
||||
return (normalized_tensor,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"OpenDiTSampler": OpenDiTSampler,
|
||||
"OpenDiTConditioning": OpenDiTConditioning,
|
||||
"DownloadAndLoadOpenSoraModel": DownloadAndLoadOpenSoraModel,
|
||||
"DownloadAndLoadOpenSoraVAE": DownloadAndLoadOpenSoraVAE,
|
||||
"DownloadAndLoadOpenDiTT5Model": DownloadAndLoadOpenDiTT5Model
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"OpenDiTSampler": "OpenDiT Sampler",
|
||||
"OpenDiTConditioning": "OpenDiT Conditioning",
|
||||
"DownloadAndLoadOpenSoraModel": "(Down)Load OpenSora Model",
|
||||
"DownloadAndLoadOpenSoraVAE": "(Down)Load OpenSora VAE",
|
||||
"DownloadAndLoadOpenDiTT5Model": "(Down)Load OpenDiT T5 Model"
|
||||
}
|
||||
Executable
Executable
Executable
+420
@@ -0,0 +1,420 @@
|
||||
from typing import Any, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from torch import Tensor
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
#from opendit.core.parallel_mgr import get_sequence_parallel_size
|
||||
|
||||
# ======================================================
|
||||
# Model
|
||||
# ======================================================
|
||||
|
||||
|
||||
def model_sharding(model: torch.nn.Module):
|
||||
global_rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
for _, param in model.named_parameters():
|
||||
padding_size = (world_size - param.numel() % world_size) % world_size
|
||||
if padding_size > 0:
|
||||
padding_param = torch.nn.functional.pad(param.data.view(-1), [0, padding_size])
|
||||
else:
|
||||
padding_param = param.data.view(-1)
|
||||
splited_params = padding_param.split(padding_param.numel() // world_size)
|
||||
splited_params = splited_params[global_rank]
|
||||
param.data = splited_params
|
||||
|
||||
|
||||
# ======================================================
|
||||
# AllGather & ReduceScatter
|
||||
# ======================================================
|
||||
|
||||
|
||||
class AsyncAllGatherForTwo(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx: Any,
|
||||
inputs: Tensor,
|
||||
weight: Tensor,
|
||||
bias: Tensor,
|
||||
sp_rank: int,
|
||||
sp_size: int,
|
||||
group: Optional[ProcessGroup] = None,
|
||||
) -> Tuple[Tensor, Any]:
|
||||
"""
|
||||
Returns:
|
||||
outputs: Tensor
|
||||
handle: Optional[Work], if overlap is True
|
||||
"""
|
||||
from torch.distributed._functional_collectives import all_gather_tensor
|
||||
|
||||
ctx.group = group
|
||||
ctx.sp_rank = sp_rank
|
||||
ctx.sp_size = sp_size
|
||||
|
||||
# all gather inputs
|
||||
all_inputs = all_gather_tensor(inputs.unsqueeze(0), 0, group)
|
||||
# compute local qkv
|
||||
local_qkv = F.linear(inputs, weight, bias).unsqueeze(0)
|
||||
|
||||
# remote compute
|
||||
remote_inputs = all_inputs[1 - sp_rank].view(list(local_qkv.shape[:-1]) + [-1])
|
||||
# compute remote qkv
|
||||
remote_qkv = F.linear(remote_inputs, weight, bias)
|
||||
|
||||
# concat local and remote qkv
|
||||
if sp_rank == 0:
|
||||
qkv = torch.cat([local_qkv, remote_qkv], dim=0)
|
||||
else:
|
||||
qkv = torch.cat([remote_qkv, local_qkv], dim=0)
|
||||
qkv = rearrange(qkv, "sp b n c -> b (sp n) c")
|
||||
|
||||
ctx.save_for_backward(inputs, weight, remote_inputs)
|
||||
return qkv
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx: Any, *grad_outputs) -> Tuple[Tensor, None, None]:
|
||||
from torch.distributed._functional_collectives import reduce_scatter_tensor
|
||||
|
||||
group = ctx.group
|
||||
sp_rank = ctx.sp_rank
|
||||
sp_size = ctx.sp_size
|
||||
inputs, weight, remote_inputs = ctx.saved_tensors
|
||||
|
||||
# split qkv_grad
|
||||
qkv_grad = grad_outputs[0]
|
||||
qkv_grad = rearrange(qkv_grad, "b (sp n) c -> sp b n c", sp=sp_size)
|
||||
qkv_grad = torch.chunk(qkv_grad, 2, dim=0)
|
||||
if sp_rank == 0:
|
||||
local_qkv_grad, remote_qkv_grad = qkv_grad
|
||||
else:
|
||||
remote_qkv_grad, local_qkv_grad = qkv_grad
|
||||
|
||||
# compute remote grad
|
||||
remote_inputs_grad = torch.matmul(remote_qkv_grad, weight).squeeze(0)
|
||||
weight_grad = torch.matmul(remote_qkv_grad.transpose(-1, -2), remote_inputs).squeeze(0).sum(0)
|
||||
bias_grad = remote_qkv_grad.squeeze(0).sum(0).sum(0)
|
||||
|
||||
# launch async reduce scatter
|
||||
remote_inputs_grad_zero = torch.zeros_like(remote_inputs_grad)
|
||||
if sp_rank == 0:
|
||||
remote_inputs_grad = torch.cat([remote_inputs_grad_zero, remote_inputs_grad], dim=0)
|
||||
else:
|
||||
remote_inputs_grad = torch.cat([remote_inputs_grad, remote_inputs_grad_zero], dim=0)
|
||||
remote_inputs_grad = reduce_scatter_tensor(remote_inputs_grad, "sum", 0, group)
|
||||
|
||||
# compute local grad and wait for reduce scatter
|
||||
local_input_grad = torch.matmul(local_qkv_grad, weight).squeeze(0)
|
||||
weight_grad += torch.matmul(local_qkv_grad.transpose(-1, -2), inputs).squeeze(0).sum(0)
|
||||
bias_grad += local_qkv_grad.squeeze(0).sum(0).sum(0)
|
||||
|
||||
# sum remote and local grad
|
||||
inputs_grad = remote_inputs_grad + local_input_grad
|
||||
return inputs_grad, weight_grad, bias_grad, None, None, None
|
||||
|
||||
|
||||
class AllGather(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx: Any,
|
||||
inputs: Tensor,
|
||||
group: Optional[ProcessGroup] = None,
|
||||
overlap: bool = False,
|
||||
) -> Tuple[Tensor, Any]:
|
||||
"""
|
||||
Returns:
|
||||
outputs: Tensor
|
||||
handle: Optional[Work], if overlap is True
|
||||
"""
|
||||
assert ctx is not None or not overlap
|
||||
|
||||
if ctx is not None:
|
||||
ctx.comm_grp = group
|
||||
|
||||
comm_size = dist.get_world_size(group)
|
||||
if comm_size == 1:
|
||||
return inputs.unsqueeze(0), None
|
||||
|
||||
buffer_shape = (comm_size,) + inputs.shape
|
||||
outputs = torch.empty(buffer_shape, dtype=inputs.dtype, device=inputs.device)
|
||||
buffer_list = list(torch.chunk(outputs, comm_size, dim=0))
|
||||
if not overlap:
|
||||
dist.all_gather(buffer_list, inputs, group=group)
|
||||
return outputs, None
|
||||
else:
|
||||
handle = dist.all_gather(buffer_list, inputs, group=group, async_op=True)
|
||||
return outputs, handle
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx: Any, *grad_outputs) -> Tuple[Tensor, None, None]:
|
||||
return (
|
||||
ReduceScatter.forward(None, grad_outputs[0], ctx.comm_grp, False)[0],
|
||||
None,
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
class ReduceScatter(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx: Any,
|
||||
inputs: Tensor,
|
||||
group: ProcessGroup,
|
||||
overlap: bool = False,
|
||||
) -> Tuple[Tensor, Any]:
|
||||
"""
|
||||
Returns:
|
||||
outputs: Tensor
|
||||
handle: Optional[Work], if overlap is True
|
||||
"""
|
||||
assert ctx is not None or not overlap
|
||||
|
||||
if ctx is not None:
|
||||
ctx.comm_grp = group
|
||||
|
||||
comm_size = dist.get_world_size(group)
|
||||
if comm_size == 1:
|
||||
return inputs.squeeze(0), None
|
||||
|
||||
if not inputs.is_contiguous():
|
||||
inputs = inputs.contiguous()
|
||||
|
||||
output_shape = inputs.shape[1:]
|
||||
outputs = torch.empty(output_shape, dtype=inputs.dtype, device=inputs.device)
|
||||
buffer_list = list(torch.chunk(inputs, comm_size, dim=0))
|
||||
if not overlap:
|
||||
dist.reduce_scatter(outputs, buffer_list, group=group)
|
||||
return outputs, None
|
||||
else:
|
||||
handle = dist.reduce_scatter(outputs, buffer_list, group=group, async_op=True)
|
||||
return outputs, handle
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx: Any, *grad_outputs) -> Tuple[Tensor, None, None]:
|
||||
# TODO: support async backward
|
||||
return (
|
||||
AllGather.forward(None, grad_outputs[0], ctx.comm_grp, False)[0],
|
||||
None,
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
# ======================================================
|
||||
# AlltoAll
|
||||
# ======================================================
|
||||
|
||||
|
||||
def _all_to_all_func(input_, world_size, group, scatter_dim, gather_dim):
|
||||
input_list = [t.contiguous() for t in torch.tensor_split(input_, world_size, scatter_dim)]
|
||||
output_list = [torch.empty_like(input_list[0]) for _ in range(world_size)]
|
||||
dist.all_to_all(output_list, input_list, group=group)
|
||||
return torch.cat(output_list, dim=gather_dim).contiguous()
|
||||
|
||||
|
||||
class _AllToAll(torch.autograd.Function):
|
||||
"""All-to-all communication.
|
||||
|
||||
Args:
|
||||
input_: input matrix
|
||||
process_group: communication group
|
||||
scatter_dim: scatter dimension
|
||||
gather_dim: gather dimension
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input_, process_group, scatter_dim, gather_dim):
|
||||
ctx.process_group = process_group
|
||||
ctx.scatter_dim = scatter_dim
|
||||
ctx.gather_dim = gather_dim
|
||||
world_size = dist.get_world_size(process_group)
|
||||
|
||||
return _all_to_all_func(input_, world_size, process_group, scatter_dim, gather_dim)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, *grad_output):
|
||||
process_group = ctx.process_group
|
||||
scatter_dim = ctx.gather_dim
|
||||
gather_dim = ctx.scatter_dim
|
||||
return_grad = _AllToAll.apply(*grad_output, process_group, scatter_dim, gather_dim)
|
||||
return (return_grad, None, None, None)
|
||||
|
||||
|
||||
def all_to_all_comm(input_, process_group=None, scatter_dim=2, gather_dim=1):
|
||||
return _AllToAll.apply(input_, process_group, scatter_dim, gather_dim)
|
||||
|
||||
|
||||
# ======================================================
|
||||
# Sequence Gather & Split
|
||||
# ======================================================
|
||||
|
||||
|
||||
def _split_sequence_func(input_, pg: dist.ProcessGroup, dim: int, pad: int):
|
||||
# skip if only one rank involved
|
||||
world_size = dist.get_world_size(pg)
|
||||
rank = dist.get_rank(pg)
|
||||
if world_size == 1:
|
||||
return input_
|
||||
|
||||
if pad > 0:
|
||||
pad_size = list(input_.shape)
|
||||
pad_size[dim] = pad
|
||||
input_ = torch.cat([input_, torch.zeros(pad_size, dtype=input_.dtype, device=input_.device)], dim=dim)
|
||||
|
||||
dim_size = input_.size(dim)
|
||||
assert dim_size % world_size == 0, f"dim_size ({dim_size}) is not divisible by world_size ({world_size})"
|
||||
|
||||
tensor_list = torch.split(input_, dim_size // world_size, dim=dim)
|
||||
output = tensor_list[rank].contiguous()
|
||||
return output
|
||||
|
||||
|
||||
def _gather_sequence_func(input_, pg: dist.ProcessGroup, dim: int, pad: int):
|
||||
# skip if only one rank involved
|
||||
input_ = input_.contiguous()
|
||||
world_size = dist.get_world_size(pg)
|
||||
dist.get_rank(pg)
|
||||
|
||||
if world_size == 1:
|
||||
return input_
|
||||
|
||||
# all gather
|
||||
tensor_list = [torch.empty_like(input_) for _ in range(world_size)]
|
||||
assert input_.device.type == "cuda"
|
||||
torch.distributed.all_gather(tensor_list, input_, group=pg)
|
||||
|
||||
# concat
|
||||
output = torch.cat(tensor_list, dim=dim)
|
||||
|
||||
if pad > 0:
|
||||
output = output.narrow(dim, 0, output.size(dim) - pad)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class _GatherForwardSplitBackward(torch.autograd.Function):
|
||||
"""
|
||||
Gather the input sequence.
|
||||
|
||||
Args:
|
||||
input_: input matrix.
|
||||
process_group: process group.
|
||||
dim: dimension
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def symbolic(graph, input_):
|
||||
return _gather_sequence_func(input_)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input_, process_group, dim, grad_scale, pad):
|
||||
ctx.process_group = process_group
|
||||
ctx.dim = dim
|
||||
ctx.grad_scale = grad_scale
|
||||
ctx.pad = pad
|
||||
return _gather_sequence_func(input_, process_group, dim, pad)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
if ctx.grad_scale == "up":
|
||||
grad_output = grad_output * dist.get_world_size(ctx.process_group)
|
||||
elif ctx.grad_scale == "down":
|
||||
grad_output = grad_output / dist.get_world_size(ctx.process_group)
|
||||
|
||||
return _split_sequence_func(grad_output, ctx.process_group, ctx.dim, ctx.pad), None, None, None, None
|
||||
|
||||
|
||||
class _SplitForwardGatherBackward(torch.autograd.Function):
|
||||
"""
|
||||
Split sequence.
|
||||
|
||||
Args:
|
||||
input_: input matrix.
|
||||
process_group: parallel mode.
|
||||
dim: dimension
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def symbolic(graph, input_):
|
||||
return _split_sequence_func(input_)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input_, process_group, dim, grad_scale, pad):
|
||||
ctx.process_group = process_group
|
||||
ctx.dim = dim
|
||||
ctx.grad_scale = grad_scale
|
||||
ctx.pad = pad
|
||||
return _split_sequence_func(input_, process_group, dim, pad)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
if ctx.grad_scale == "up":
|
||||
grad_output = grad_output * dist.get_world_size(ctx.process_group)
|
||||
elif ctx.grad_scale == "down":
|
||||
grad_output = grad_output / dist.get_world_size(ctx.process_group)
|
||||
return _gather_sequence_func(grad_output, ctx.process_group, ctx.pad), None, None, None, None
|
||||
|
||||
|
||||
def split_sequence(input_, process_group, dim, grad_scale=1.0, pad=0):
|
||||
return _SplitForwardGatherBackward.apply(input_, process_group, dim, grad_scale, pad)
|
||||
|
||||
|
||||
def gather_sequence(input_, process_group, dim, grad_scale=1.0, pad=0):
|
||||
return _GatherForwardSplitBackward.apply(input_, process_group, dim, grad_scale, pad)
|
||||
|
||||
|
||||
# ==============================
|
||||
# Pad
|
||||
# ==============================
|
||||
|
||||
SPTIAL_PAD = 0
|
||||
TEMPORAL_PAD = 0
|
||||
|
||||
|
||||
def set_spatial_pad(dim_size: int):
|
||||
sp_size = get_sequence_parallel_size()
|
||||
pad = (sp_size - (dim_size % sp_size)) % sp_size
|
||||
global SPTIAL_PAD
|
||||
SPTIAL_PAD = pad
|
||||
|
||||
|
||||
def get_spatial_pad() -> int:
|
||||
return SPTIAL_PAD
|
||||
|
||||
|
||||
def set_temporal_pad(dim_size: int):
|
||||
sp_size = get_sequence_parallel_size()
|
||||
pad = (sp_size - (dim_size % sp_size)) % sp_size
|
||||
global TEMPORAL_PAD
|
||||
TEMPORAL_PAD = pad
|
||||
|
||||
|
||||
def get_temporal_pad() -> int:
|
||||
return TEMPORAL_PAD
|
||||
|
||||
|
||||
def all_to_all_with_pad(
|
||||
input_: torch.Tensor,
|
||||
process_group: dist.ProcessGroup,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1,
|
||||
scatter_pad: int = 0,
|
||||
gather_pad: int = 0,
|
||||
):
|
||||
if scatter_pad > 0:
|
||||
pad_shape = list(input_.shape)
|
||||
pad_shape[scatter_dim] = scatter_pad
|
||||
pad_tensor = torch.zeros(pad_shape, device=input_.device, dtype=input_.dtype)
|
||||
input_ = torch.cat([input_, pad_tensor], dim=scatter_dim)
|
||||
|
||||
assert (
|
||||
input_.shape[scatter_dim] % dist.get_world_size(process_group) == 0
|
||||
), f"Dimension to scatter ({input_.shape[scatter_dim]}) is not divisible by world size ({dist.get_world_size(process_group)})"
|
||||
input_ = _AllToAll.apply(input_, process_group, scatter_dim, gather_dim)
|
||||
|
||||
if gather_pad > 0:
|
||||
input_ = input_.narrow(gather_dim, 0, input_.size(gather_dim) - gather_pad)
|
||||
|
||||
return input_
|
||||
Executable
+245
@@ -0,0 +1,245 @@
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
#import torch.distributed as dist
|
||||
|
||||
PAB_MANAGER = None
|
||||
|
||||
|
||||
class PABManager:
|
||||
def __init__(
|
||||
self,
|
||||
steps: int = 100,
|
||||
cross_broadcast: bool = False,
|
||||
cross_threshold: list = [100, 900],
|
||||
cross_gap: int = 5,
|
||||
spatial_broadcast: bool = False,
|
||||
spatial_threshold: list = [100, 900],
|
||||
spatial_gap: int = 2,
|
||||
temporal_broadcast: bool = False,
|
||||
temporal_threshold: list = [100, 900],
|
||||
temporal_gap: int = 3,
|
||||
diffusion_skip: bool = False,
|
||||
diffusion_timestep_respacing: list = None,
|
||||
diffusion_skip_timestep: list = None,
|
||||
):
|
||||
self.steps = steps
|
||||
|
||||
self.cross_broadcast = cross_broadcast
|
||||
self.cross_threshold = cross_threshold
|
||||
self.cross_gap = cross_gap
|
||||
|
||||
self.spatial_broadcast = spatial_broadcast
|
||||
self.spatial_threshold = spatial_threshold
|
||||
self.spatial_gap = spatial_gap
|
||||
|
||||
self.temporal_broadcast = temporal_broadcast
|
||||
self.temporal_threshold = temporal_threshold
|
||||
self.temporal_gap = temporal_gap
|
||||
|
||||
self.diffusion_skip = diffusion_skip
|
||||
self.diffusion_timestep_respacing = diffusion_timestep_respacing
|
||||
self.diffusion_skip_timestep = diffusion_skip_timestep
|
||||
|
||||
|
||||
print(
|
||||
f"\n\
|
||||
Init SkipManager:\n\
|
||||
steps={steps}\n\
|
||||
cross_broadcast={cross_broadcast}, cross_threshold={cross_threshold}, cross_gap={cross_gap}\n\
|
||||
spatial_broadcast={spatial_broadcast}, spatial_threshold={spatial_threshold}, spatial_gap={spatial_gap}\n\
|
||||
temporal_broadcast={temporal_broadcast}, temporal_threshold={temporal_threshold}, temporal_gap={temporal_gap}\n\
|
||||
\n",
|
||||
end="",
|
||||
)
|
||||
|
||||
def if_broadcast_cross(self, timestep: int, count: int):
|
||||
if (
|
||||
self.cross_broadcast
|
||||
and (timestep is not None)
|
||||
and (count % self.cross_gap != 0)
|
||||
and (self.cross_threshold[0] < timestep < self.cross_threshold[1])
|
||||
):
|
||||
flag = True
|
||||
else:
|
||||
flag = False
|
||||
count = (count + 1) % self.steps
|
||||
return flag, count
|
||||
|
||||
def if_broadcast_temporal(self, timestep: int, count: int):
|
||||
if (
|
||||
self.temporal_broadcast
|
||||
and (timestep is not None)
|
||||
and (count % self.temporal_gap != 0)
|
||||
and (self.temporal_threshold[0] < timestep < self.temporal_threshold[1])
|
||||
):
|
||||
flag = True
|
||||
else:
|
||||
flag = False
|
||||
count = (count + 1) % self.steps
|
||||
return flag, count
|
||||
|
||||
def if_broadcast_spatial(self, timestep: int, count: int, block_idx: int):
|
||||
if (
|
||||
self.spatial_broadcast
|
||||
and (timestep is not None)
|
||||
and (count % self.spatial_gap != 0)
|
||||
and (self.spatial_threshold[0] < timestep < self.spatial_threshold[1])
|
||||
):
|
||||
flag = True
|
||||
else:
|
||||
flag = False
|
||||
count = (count + 1) % self.steps
|
||||
return flag, count
|
||||
|
||||
|
||||
def set_pab_manager(
|
||||
steps: int = 100,
|
||||
cross_broadcast: bool = False,
|
||||
cross_threshold: list = [100, 900],
|
||||
cross_gap: int = 5,
|
||||
spatial_broadcast: bool = False,
|
||||
spatial_threshold: list = [100, 900],
|
||||
spatial_gap: int = 2,
|
||||
temporal_broadcast: bool = False,
|
||||
temporal_threshold: list = [100, 900],
|
||||
temporal_gap: int = 3,
|
||||
diffusion_skip: bool = False,
|
||||
diffusion_timestep_respacing: list = None,
|
||||
diffusion_skip_timestep: list = None,
|
||||
):
|
||||
global PAB_MANAGER
|
||||
PAB_MANAGER = PABManager(
|
||||
steps,
|
||||
cross_broadcast,
|
||||
cross_threshold,
|
||||
cross_gap,
|
||||
spatial_broadcast,
|
||||
spatial_threshold,
|
||||
spatial_gap,
|
||||
temporal_broadcast,
|
||||
temporal_threshold,
|
||||
temporal_gap,
|
||||
diffusion_skip,
|
||||
diffusion_timestep_respacing,
|
||||
diffusion_skip_timestep,
|
||||
)
|
||||
|
||||
|
||||
def enable_pab():
|
||||
if PAB_MANAGER is None:
|
||||
return False
|
||||
return PAB_MANAGER.cross_broadcast or PAB_MANAGER.spatial_broadcast or PAB_MANAGER.temporal_broadcast
|
||||
|
||||
|
||||
def if_broadcast_cross(timestep: int, count: int):
|
||||
if not enable_pab():
|
||||
return False, count
|
||||
return PAB_MANAGER.if_broadcast_cross(timestep, count)
|
||||
|
||||
|
||||
def if_broadcast_temporal(timestep: int, count: int):
|
||||
if not enable_pab():
|
||||
return False, count
|
||||
return PAB_MANAGER.if_broadcast_temporal(timestep, count)
|
||||
|
||||
|
||||
def if_broadcast_spatial(timestep: int, count: int, block_idx: int):
|
||||
if not enable_pab():
|
||||
return False, count
|
||||
return PAB_MANAGER.if_broadcast_spatial(timestep, count, block_idx)
|
||||
|
||||
|
||||
def get_diffusion_skip():
|
||||
return enable_pab() and PAB_MANAGER.diffusion_skip
|
||||
|
||||
|
||||
def get_diffusion_timestep_respacing():
|
||||
return PAB_MANAGER.diffusion_timestep_respacing
|
||||
|
||||
|
||||
def get_diffusion_skip_timestep():
|
||||
return enable_pab() and PAB_MANAGER.diffusion_skip_timestep
|
||||
|
||||
|
||||
def space_timesteps(time_steps, time_bins):
|
||||
num_bins = len(time_bins)
|
||||
bin_size = time_steps // num_bins
|
||||
|
||||
result = []
|
||||
|
||||
for i, bin_count in enumerate(time_bins):
|
||||
start = i * bin_size
|
||||
end = start + bin_size
|
||||
|
||||
bin_steps = np.linspace(start, end, bin_count, endpoint=False, dtype=int).tolist()
|
||||
result.extend(bin_steps)
|
||||
|
||||
result_tensor = torch.tensor(result, dtype=torch.int32)
|
||||
sorted_tensor = torch.sort(result_tensor, descending=True).values
|
||||
|
||||
return sorted_tensor
|
||||
|
||||
|
||||
def skip_diffusion_timestep(timesteps, diffusion_skip_timestep):
|
||||
if isinstance(timesteps, list):
|
||||
# If timesteps is a list, we assume each element is a tensor
|
||||
timesteps_np = [t.cpu().numpy() for t in timesteps]
|
||||
device = timesteps[0].device
|
||||
else:
|
||||
# If timesteps is a tensor
|
||||
timesteps_np = timesteps.cpu().numpy()
|
||||
device = timesteps.device
|
||||
|
||||
num_bins = len(diffusion_skip_timestep)
|
||||
|
||||
if isinstance(timesteps_np, list):
|
||||
bin_size = len(timesteps_np) // num_bins
|
||||
new_timesteps = []
|
||||
|
||||
for i in range(num_bins):
|
||||
bin_start = i * bin_size
|
||||
bin_end = (i + 1) * bin_size if i != num_bins - 1 else len(timesteps_np)
|
||||
bin_timesteps = timesteps_np[bin_start:bin_end]
|
||||
|
||||
if diffusion_skip_timestep[i] == 0:
|
||||
# If the bin is marked with 0, keep all timesteps
|
||||
new_timesteps.extend(bin_timesteps)
|
||||
elif diffusion_skip_timestep[i] == 1:
|
||||
# If the bin is marked with 1, omit the last timestep in the bin
|
||||
new_timesteps.extend(bin_timesteps[1:])
|
||||
|
||||
new_timesteps_tensor = [torch.tensor(t, device=device) for t in new_timesteps]
|
||||
else:
|
||||
bin_size = len(timesteps_np) // num_bins
|
||||
new_timesteps = []
|
||||
|
||||
for i in range(num_bins):
|
||||
bin_start = i * bin_size
|
||||
bin_end = (i + 1) * bin_size if i != num_bins - 1 else len(timesteps_np)
|
||||
bin_timesteps = timesteps_np[bin_start:bin_end]
|
||||
|
||||
if diffusion_skip_timestep[i] == 0:
|
||||
# If the bin is marked with 0, keep all timesteps
|
||||
new_timesteps.extend(bin_timesteps)
|
||||
elif diffusion_skip_timestep[i] == 1:
|
||||
# If the bin is marked with 1, omit the last timestep in the bin
|
||||
new_timesteps.extend(bin_timesteps[1:])
|
||||
elif diffusion_skip_timestep[i] != 0:
|
||||
# If the bin is marked with a non-zero value, randomly omit n timesteps
|
||||
if len(bin_timesteps) > diffusion_skip_timestep[i]:
|
||||
indices_to_remove = set(random.sample(range(len(bin_timesteps)), diffusion_skip_timestep[i]))
|
||||
timesteps_to_keep = [
|
||||
timestep for idx, timestep in enumerate(bin_timesteps) if idx not in indices_to_remove
|
||||
]
|
||||
else:
|
||||
timesteps_to_keep = bin_timesteps # 如果bin_timesteps的长度小于等于n,则不删除任何元素
|
||||
new_timesteps.extend(timesteps_to_keep)
|
||||
|
||||
new_timesteps_tensor = torch.tensor(new_timesteps, device=device)
|
||||
|
||||
if isinstance(timesteps, list):
|
||||
return new_timesteps_tensor
|
||||
else:
|
||||
return new_timesteps_tensor
|
||||
Executable
+54
@@ -0,0 +1,54 @@
|
||||
import torch.distributed as dist
|
||||
from colossalai.cluster.process_group_mesh import ProcessGroupMesh
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
PARALLEL_MANAGER = None
|
||||
|
||||
|
||||
class ParallelManager(ProcessGroupMesh):
|
||||
def __init__(self, dp_size, sp_size, dp_axis, sp_axis):
|
||||
super().__init__(dp_size, sp_size)
|
||||
self.dp_axis = dp_axis
|
||||
self.dp_group: ProcessGroup = self.get_group_along_axis(self.dp_axis)
|
||||
self.dp_rank = dist.get_rank(self.dp_group)
|
||||
|
||||
self.sp_size = sp_size
|
||||
self.sp_axis = sp_axis
|
||||
self.sp_group: ProcessGroup = self.get_group_along_axis(self.sp_axis)
|
||||
self.sp_rank = dist.get_rank(self.sp_group)
|
||||
self.enable_sp = sp_size > 1
|
||||
|
||||
|
||||
def set_parallel_manager(dp_size, sp_size, dp_axis=0, sp_axis=1):
|
||||
global PARALLEL_MANAGER
|
||||
PARALLEL_MANAGER = ParallelManager(dp_size, sp_size, dp_axis, sp_axis)
|
||||
|
||||
|
||||
def get_data_parallel_group():
|
||||
return PARALLEL_MANAGER.dp_group
|
||||
|
||||
|
||||
def get_data_parallel_rank():
|
||||
return PARALLEL_MANAGER.dp_rank
|
||||
|
||||
|
||||
def get_sequence_parallel_group():
|
||||
return PARALLEL_MANAGER.sp_group
|
||||
|
||||
|
||||
def get_sequence_parallel_size():
|
||||
return PARALLEL_MANAGER.sp_size
|
||||
|
||||
|
||||
def get_sequence_parallel_rank():
|
||||
return PARALLEL_MANAGER.sp_rank
|
||||
|
||||
|
||||
def enable_sequence_parallel():
|
||||
if PARALLEL_MANAGER is None:
|
||||
return False
|
||||
return PARALLEL_MANAGER.enable_sp
|
||||
|
||||
|
||||
def get_parallel_manager():
|
||||
return PARALLEL_MANAGER
|
||||
Executable
Executable
Executable
+39
@@ -0,0 +1,39 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class T5LayerNorm(nn.Module):
|
||||
def __init__(self, hidden_size, eps=1e-6):
|
||||
"""
|
||||
Construct a layernorm module in the T5 style. No bias and no subtraction of mean.
|
||||
"""
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(hidden_size))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
def forward(self, hidden_states):
|
||||
# T5 uses a layer_norm which only scales and doesn't shift, which is also known as Root Mean
|
||||
# Square Layer Normalization https://arxiv.org/abs/1910.07467 thus varience is calculated
|
||||
# w/o mean and there is no bias. Additionally we want to make sure that the accumulation for
|
||||
# half-precision inputs is done in fp32
|
||||
|
||||
variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
|
||||
|
||||
# convert into half-precision if necessary
|
||||
if self.weight.dtype in [torch.float16, torch.bfloat16]:
|
||||
hidden_states = hidden_states.to(self.weight.dtype)
|
||||
|
||||
return self.weight * hidden_states
|
||||
|
||||
@staticmethod
|
||||
def from_native_module(module, *args, **kwargs):
|
||||
assert module.__class__.__name__ == "FusedRMSNorm", (
|
||||
"Recovering T5LayerNorm requires the original layer to be apex's Fused RMS Norm."
|
||||
"Apex's fused norm is automatically used by Hugging Face Transformers https://github.com/huggingface/transformers/blob/main/src/transformers/models/t5/modeling_t5.py#L265C5-L265C48"
|
||||
)
|
||||
|
||||
layer_norm = T5LayerNorm(module.normalized_shape, eps=module.eps)
|
||||
layer_norm.weight.data.copy_(module.weight.data)
|
||||
layer_norm = layer_norm.to(module.weight.device)
|
||||
return layer_norm
|
||||
Executable
+68
@@ -0,0 +1,68 @@
|
||||
from colossalai.shardformer.modeling.jit import get_jit_fused_dropout_add_func
|
||||
from colossalai.shardformer.modeling.t5 import get_jit_fused_T5_layer_ff_forward, get_T5_layer_self_attention_forward
|
||||
from colossalai.shardformer.policies.base_policy import Policy, SubModuleReplacementDescription
|
||||
|
||||
|
||||
class T5EncoderPolicy(Policy):
|
||||
def config_sanity_check(self):
|
||||
assert not self.shard_config.enable_tensor_parallelism
|
||||
assert not self.shard_config.enable_flash_attention
|
||||
|
||||
def preprocess(self):
|
||||
return self.model
|
||||
|
||||
def module_policy(self):
|
||||
from transformers.models.t5.modeling_t5 import T5LayerFF, T5LayerSelfAttention, T5Stack
|
||||
|
||||
policy = {}
|
||||
|
||||
# check whether apex is installed
|
||||
try:
|
||||
from apex.normalization import FusedRMSNorm # noqa
|
||||
from opendit.core.shardformer.t5.modeling import T5LayerNorm
|
||||
|
||||
# recover hf from fused rms norm to T5 norm which is faster
|
||||
self.append_or_create_submodule_replacement(
|
||||
description=SubModuleReplacementDescription(
|
||||
suffix="layer_norm",
|
||||
target_module=T5LayerNorm,
|
||||
),
|
||||
policy=policy,
|
||||
target_key=T5LayerFF,
|
||||
)
|
||||
self.append_or_create_submodule_replacement(
|
||||
description=SubModuleReplacementDescription(suffix="layer_norm", target_module=T5LayerNorm),
|
||||
policy=policy,
|
||||
target_key=T5LayerSelfAttention,
|
||||
)
|
||||
self.append_or_create_submodule_replacement(
|
||||
description=SubModuleReplacementDescription(suffix="final_layer_norm", target_module=T5LayerNorm),
|
||||
policy=policy,
|
||||
target_key=T5Stack,
|
||||
)
|
||||
except (ImportError, ModuleNotFoundError):
|
||||
pass
|
||||
|
||||
# use jit operator
|
||||
if self.shard_config.enable_jit_fused:
|
||||
self.append_or_create_method_replacement(
|
||||
description={
|
||||
"forward": get_jit_fused_T5_layer_ff_forward(),
|
||||
"dropout_add": get_jit_fused_dropout_add_func(),
|
||||
},
|
||||
policy=policy,
|
||||
target_key=T5LayerFF,
|
||||
)
|
||||
self.append_or_create_method_replacement(
|
||||
description={
|
||||
"forward": get_T5_layer_self_attention_forward(),
|
||||
"dropout_add": get_jit_fused_dropout_add_func(),
|
||||
},
|
||||
policy=policy,
|
||||
target_key=T5LayerSelfAttention,
|
||||
)
|
||||
|
||||
return policy
|
||||
|
||||
def postprocess(self):
|
||||
return self.model
|
||||
Executable
+94
@@ -0,0 +1,94 @@
|
||||
import random
|
||||
from typing import Iterator, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.utils.data import DataLoader, Dataset, DistributedSampler
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
|
||||
from opendit.core.parallel_mgr import ParallelManager
|
||||
|
||||
|
||||
class StatefulDistributedSampler(DistributedSampler):
|
||||
def __init__(
|
||||
self,
|
||||
dataset: Dataset,
|
||||
num_replicas: Optional[int] = None,
|
||||
rank: Optional[int] = None,
|
||||
shuffle: bool = True,
|
||||
seed: int = 0,
|
||||
drop_last: bool = False,
|
||||
) -> None:
|
||||
super().__init__(dataset, num_replicas, rank, shuffle, seed, drop_last)
|
||||
self.start_index: int = 0
|
||||
|
||||
def __iter__(self) -> Iterator:
|
||||
iterator = super().__iter__()
|
||||
indices = list(iterator)
|
||||
indices = indices[self.start_index :]
|
||||
return iter(indices)
|
||||
|
||||
def __len__(self) -> int:
|
||||
return self.num_samples - self.start_index
|
||||
|
||||
def set_start_index(self, start_index: int) -> None:
|
||||
self.start_index = start_index
|
||||
|
||||
|
||||
def prepare_dataloader(
|
||||
dataset,
|
||||
batch_size,
|
||||
shuffle=False,
|
||||
seed=1024,
|
||||
drop_last=False,
|
||||
pin_memory=False,
|
||||
num_workers=0,
|
||||
pg_manager: Optional[ParallelManager] = None,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
Prepare a dataloader for distributed training. The dataloader will be wrapped by
|
||||
`torch.utils.data.DataLoader` and `StatefulDistributedSampler`.
|
||||
|
||||
|
||||
Args:
|
||||
dataset (`torch.utils.data.Dataset`): The dataset to be loaded.
|
||||
shuffle (bool, optional): Whether to shuffle the dataset. Defaults to False.
|
||||
seed (int, optional): Random worker seed for sampling, defaults to 1024.
|
||||
add_sampler: Whether to add ``DistributedDataParallelSampler`` to the dataset. Defaults to True.
|
||||
drop_last (bool, optional): Set to True to drop the last incomplete batch, if the dataset size
|
||||
is not divisible by the batch size. If False and the size of dataset is not divisible by
|
||||
the batch size, then the last batch will be smaller, defaults to False.
|
||||
pin_memory (bool, optional): Whether to pin memory address in CPU memory. Defaults to False.
|
||||
num_workers (int, optional): Number of worker threads for this dataloader. Defaults to 0.
|
||||
kwargs (dict): optional parameters for ``torch.utils.data.DataLoader``, more details could be found in
|
||||
`DataLoader <https://pytorch.org/docs/stable/_modules/torch/utils/data/dataloader.html#DataLoader>`_.
|
||||
|
||||
Returns:
|
||||
:class:`torch.utils.data.DataLoader`: A DataLoader used for training or testing.
|
||||
"""
|
||||
_kwargs = kwargs.copy()
|
||||
sampler = StatefulDistributedSampler(
|
||||
dataset,
|
||||
num_replicas=pg_manager.size(pg_manager.dp_axis),
|
||||
rank=pg_manager.coordinate(pg_manager.dp_axis),
|
||||
shuffle=shuffle,
|
||||
)
|
||||
|
||||
# Deterministic dataloader
|
||||
def seed_worker(worker_id):
|
||||
worker_seed = seed
|
||||
np.random.seed(worker_seed)
|
||||
torch.manual_seed(worker_seed)
|
||||
random.seed(worker_seed)
|
||||
|
||||
return DataLoader(
|
||||
dataset,
|
||||
batch_size=batch_size,
|
||||
sampler=sampler,
|
||||
worker_init_fn=seed_worker,
|
||||
drop_last=drop_last,
|
||||
pin_memory=pin_memory,
|
||||
num_workers=num_workers,
|
||||
**_kwargs,
|
||||
)
|
||||
Executable
+42
@@ -0,0 +1,42 @@
|
||||
# Adapted from DiT
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# DiT: https://github.com/facebookresearch/DiT
|
||||
# --------------------------------------------------------
|
||||
|
||||
|
||||
import numpy as np
|
||||
import torchvision.transforms as transforms
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def center_crop_arr(pil_image, image_size):
|
||||
"""
|
||||
Center cropping implementation from ADM.
|
||||
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
|
||||
"""
|
||||
while min(*pil_image.size) >= 2 * image_size:
|
||||
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size), resample=Image.BOX)
|
||||
|
||||
scale = image_size / min(*pil_image.size)
|
||||
pil_image = pil_image.resize(tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC)
|
||||
|
||||
arr = np.array(pil_image)
|
||||
crop_y = (arr.shape[0] - image_size) // 2
|
||||
crop_x = (arr.shape[1] - image_size) // 2
|
||||
return Image.fromarray(arr[crop_y : crop_y + image_size, crop_x : crop_x + image_size])
|
||||
|
||||
|
||||
def get_transforms_image(image_size=256):
|
||||
transform = transforms.Compose(
|
||||
[
|
||||
transforms.Lambda(lambda pil_image: center_crop_arr(pil_image, image_size)),
|
||||
transforms.RandomHorizontalFlip(),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
]
|
||||
)
|
||||
return transform
|
||||
Executable
+441
@@ -0,0 +1,441 @@
|
||||
# Adapted from OpenSora and Latte
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# OpenSora: https://github.com/hpcaitech/Open-Sora
|
||||
# Latte: https://github.com/Vchitect/Latte
|
||||
# --------------------------------------------------------
|
||||
|
||||
import numbers
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def _is_tensor_video_clip(clip):
|
||||
if not torch.is_tensor(clip):
|
||||
raise TypeError("clip should be Tensor. Got %s" % type(clip))
|
||||
|
||||
if not clip.ndimension() == 4:
|
||||
raise ValueError("clip should be 4D. Got %dD" % clip.dim())
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def center_crop_arr(pil_image, image_size):
|
||||
"""
|
||||
Center cropping implementation from ADM.
|
||||
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
|
||||
"""
|
||||
while min(*pil_image.size) >= 2 * image_size:
|
||||
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size), resample=Image.BOX)
|
||||
|
||||
scale = image_size / min(*pil_image.size)
|
||||
pil_image = pil_image.resize(tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC)
|
||||
|
||||
arr = np.array(pil_image)
|
||||
crop_y = (arr.shape[0] - image_size) // 2
|
||||
crop_x = (arr.shape[1] - image_size) // 2
|
||||
return Image.fromarray(arr[crop_y : crop_y + image_size, crop_x : crop_x + image_size])
|
||||
|
||||
|
||||
def crop(clip, i, j, h, w):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
"""
|
||||
if len(clip.size()) != 4:
|
||||
raise ValueError("clip should be a 4D tensor")
|
||||
return clip[..., i : i + h, j : j + w]
|
||||
|
||||
|
||||
def resize(clip, target_size, interpolation_mode):
|
||||
if len(target_size) != 2:
|
||||
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
|
||||
return torch.nn.functional.interpolate(clip, size=target_size, mode=interpolation_mode, align_corners=False)
|
||||
|
||||
|
||||
def resize_scale(clip, target_size, interpolation_mode):
|
||||
if len(target_size) != 2:
|
||||
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
|
||||
H, W = clip.size(-2), clip.size(-1)
|
||||
scale_ = target_size[0] / min(H, W)
|
||||
return torch.nn.functional.interpolate(clip, scale_factor=scale_, mode=interpolation_mode, align_corners=False)
|
||||
|
||||
|
||||
def resized_crop(clip, i, j, h, w, size, interpolation_mode="bilinear"):
|
||||
"""
|
||||
Do spatial cropping and resizing to the video clip
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
i (int): i in (i,j) i.e coordinates of the upper left corner.
|
||||
j (int): j in (i,j) i.e coordinates of the upper left corner.
|
||||
h (int): Height of the cropped region.
|
||||
w (int): Width of the cropped region.
|
||||
size (tuple(int, int)): height and width of resized clip
|
||||
Returns:
|
||||
clip (torch.tensor): Resized and cropped clip. Size is (T, C, H, W)
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
clip = crop(clip, i, j, h, w)
|
||||
clip = resize(clip, size, interpolation_mode)
|
||||
return clip
|
||||
|
||||
|
||||
def center_crop(clip, crop_size):
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
th, tw = crop_size
|
||||
if h < th or w < tw:
|
||||
raise ValueError("height and width must be no smaller than crop_size")
|
||||
|
||||
i = int(round((h - th) / 2.0))
|
||||
j = int(round((w - tw) / 2.0))
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
def center_crop_using_short_edge(clip):
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
if h < w:
|
||||
th, tw = h, h
|
||||
i = 0
|
||||
j = int(round((w - tw) / 2.0))
|
||||
else:
|
||||
th, tw = w, w
|
||||
i = int(round((h - th) / 2.0))
|
||||
j = 0
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
def random_shift_crop(clip):
|
||||
"""
|
||||
Slide along the long edge, with the short edge as crop size
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
|
||||
if h <= w:
|
||||
short_edge = h
|
||||
else:
|
||||
short_edge = w
|
||||
|
||||
th, tw = short_edge, short_edge
|
||||
|
||||
i = torch.randint(0, h - th + 1, size=(1,)).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1,)).item()
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
def to_tensor(clip):
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
permute the dimensions of clip tensor
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
_is_tensor_video_clip(clip)
|
||||
if not clip.dtype == torch.uint8:
|
||||
raise TypeError("clip tensor should have data type uint8. Got %s" % str(clip.dtype))
|
||||
# return clip.float().permute(3, 0, 1, 2) / 255.0
|
||||
return clip.float() / 255.0
|
||||
|
||||
|
||||
def normalize(clip, mean, std, inplace=False):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
|
||||
mean (tuple): pixel RGB mean. Size is (3)
|
||||
std (tuple): pixel standard deviation. Size is (3)
|
||||
Returns:
|
||||
normalized clip (torch.tensor): Size is (T, C, H, W)
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
if not inplace:
|
||||
clip = clip.clone()
|
||||
mean = torch.as_tensor(mean, dtype=clip.dtype, device=clip.device)
|
||||
# print(mean)
|
||||
std = torch.as_tensor(std, dtype=clip.dtype, device=clip.device)
|
||||
clip.sub_(mean[:, None, None, None]).div_(std[:, None, None, None])
|
||||
return clip
|
||||
|
||||
|
||||
def hflip(clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
|
||||
Returns:
|
||||
flipped clip (torch.tensor): Size is (T, C, H, W)
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
return clip.flip(-1)
|
||||
|
||||
|
||||
class RandomCropVideo:
|
||||
def __init__(self, size):
|
||||
if isinstance(size, numbers.Number):
|
||||
self.size = (int(size), int(size))
|
||||
else:
|
||||
self.size = size
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: randomly cropped video clip.
|
||||
size is (T, C, OH, OW)
|
||||
"""
|
||||
i, j, h, w = self.get_params(clip)
|
||||
return crop(clip, i, j, h, w)
|
||||
|
||||
def get_params(self, clip):
|
||||
h, w = clip.shape[-2:]
|
||||
th, tw = self.size
|
||||
|
||||
if h < th or w < tw:
|
||||
raise ValueError(f"Required crop size {(th, tw)} is larger than input image size {(h, w)}")
|
||||
|
||||
if w == tw and h == th:
|
||||
return 0, 0, h, w
|
||||
|
||||
i = torch.randint(0, h - th + 1, size=(1,)).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1,)).item()
|
||||
|
||||
return i, j, th, tw
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size})"
|
||||
|
||||
|
||||
class CenterCropResizeVideo:
|
||||
"""
|
||||
First use the short side for cropping length,
|
||||
center crop video, then resize to the specified size
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: scale resized / center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
clip_center_crop = center_crop_using_short_edge(clip)
|
||||
clip_center_crop_resize = resize(
|
||||
clip_center_crop, target_size=self.size, interpolation_mode=self.interpolation_mode
|
||||
)
|
||||
return clip_center_crop_resize
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class UCFCenterCropVideo:
|
||||
"""
|
||||
First scale to the specified size in equal proportion to the short edge,
|
||||
then center cropping
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: scale resized / center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
clip_resize = resize_scale(clip=clip, target_size=self.size, interpolation_mode=self.interpolation_mode)
|
||||
clip_center_crop = center_crop(clip_resize, self.size)
|
||||
return clip_center_crop
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class KineticsRandomCropResizeVideo:
|
||||
"""
|
||||
Slide along the long edge, with the short edge as crop size. And resie to the desired size.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
clip_random_crop = random_shift_crop(clip)
|
||||
clip_resize = resize(clip_random_crop, self.size, self.interpolation_mode)
|
||||
return clip_resize
|
||||
|
||||
|
||||
class CenterCropVideo:
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
clip_center_crop = center_crop(clip, self.size)
|
||||
return clip_center_crop
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class NormalizeVideo:
|
||||
"""
|
||||
Normalize the video clip by mean subtraction and division by standard deviation
|
||||
Args:
|
||||
mean (3-tuple): pixel RGB mean
|
||||
std (3-tuple): pixel RGB standard deviation
|
||||
inplace (boolean): whether do in-place normalization
|
||||
"""
|
||||
|
||||
def __init__(self, mean, std, inplace=False):
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
self.inplace = inplace
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): video clip must be normalized. Size is (C, T, H, W)
|
||||
"""
|
||||
return normalize(clip, self.mean, self.std, self.inplace)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(mean={self.mean}, std={self.std}, inplace={self.inplace})"
|
||||
|
||||
|
||||
class ToTensorVideo:
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
permute the dimensions of clip tensor
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
return to_tensor(clip)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return self.__class__.__name__
|
||||
|
||||
|
||||
class RandomHorizontalFlipVideo:
|
||||
"""
|
||||
Flip the video clip along the horizontal direction with a given probability
|
||||
Args:
|
||||
p (float): probability of the clip being flipped. Default value is 0.5
|
||||
"""
|
||||
|
||||
def __init__(self, p=0.5):
|
||||
self.p = p
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor): Size is (T, C, H, W)
|
||||
"""
|
||||
if random.random() < self.p:
|
||||
clip = hflip(clip)
|
||||
return clip
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(p={self.p})"
|
||||
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# --------------------- Sampling ---------------------------
|
||||
# ------------------------------------------------------------
|
||||
class TemporalRandomCrop(object):
|
||||
"""Temporally crop the given frame indices at a random location.
|
||||
|
||||
Args:
|
||||
size (int): Desired length of frames will be seen in the model.
|
||||
"""
|
||||
|
||||
def __init__(self, size):
|
||||
self.size = size
|
||||
|
||||
def __call__(self, total_frames):
|
||||
rand_end = max(0, total_frames - self.size - 1)
|
||||
begin_index = random.randint(0, rand_end)
|
||||
end_index = min(begin_index + self.size, total_frames)
|
||||
return begin_index, end_index
|
||||
Executable
+41
@@ -0,0 +1,41 @@
|
||||
# Modified from OpenAI's diffusion repos and Meta DiT
|
||||
# DiT: https://github.com/facebookresearch/DiT/tree/main
|
||||
# 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 create_diffusion(
|
||||
timestep_respacing,
|
||||
noise_schedule="linear",
|
||||
use_kl=False,
|
||||
sigma_small=False,
|
||||
predict_xstart=False,
|
||||
learn_sigma=True,
|
||||
rescale_learned_sigmas=False,
|
||||
diffusion_steps=1000,
|
||||
):
|
||||
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.EPSILON if not predict_xstart else gd.ModelMeanType.START_X),
|
||||
model_var_type=(
|
||||
(gd.ModelVarType.FIXED_LARGE if not sigma_small else gd.ModelVarType.FIXED_SMALL)
|
||||
if not learn_sigma
|
||||
else gd.ModelVarType.LEARNED_RANGE
|
||||
),
|
||||
loss_type=loss_type
|
||||
# rescale_timesteps=rescale_timesteps,
|
||||
)
|
||||
Executable
+79
@@ -0,0 +1,79 @@
|
||||
# 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
|
||||
|
||||
import numpy as np
|
||||
import torch as th
|
||||
|
||||
|
||||
def normal_kl(mean1, logvar1, mean2, logvar2):
|
||||
"""
|
||||
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, th.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 th.exp().
|
||||
logvar1, logvar2 = [x if isinstance(x, th.Tensor) else th.tensor(x).to(tensor) for x in (logvar1, logvar2)]
|
||||
|
||||
return 0.5 * (-1.0 + logvar2 - logvar1 + th.exp(logvar1 - logvar2) + ((mean1 - mean2) ** 2) * th.exp(-logvar2))
|
||||
|
||||
|
||||
def approx_standard_normal_cdf(x):
|
||||
"""
|
||||
A fast approximation of the cumulative distribution function of the
|
||||
standard normal.
|
||||
"""
|
||||
return 0.5 * (1.0 + th.tanh(np.sqrt(2.0 / np.pi) * (x + 0.044715 * th.pow(x, 3))))
|
||||
|
||||
|
||||
def continuous_gaussian_log_likelihood(x, *, means, log_scales):
|
||||
"""
|
||||
Compute the log-likelihood of a continuous Gaussian distribution.
|
||||
:param x: the targets
|
||||
:param means: the Gaussian mean Tensor.
|
||||
:param log_scales: the Gaussian log stddev Tensor.
|
||||
:return: a tensor like x of log probabilities (in nats).
|
||||
"""
|
||||
centered_x = x - means
|
||||
inv_stdv = th.exp(-log_scales)
|
||||
normalized_x = centered_x * inv_stdv
|
||||
log_probs = th.distributions.Normal(th.zeros_like(x), th.ones_like(x)).log_prob(normalized_x)
|
||||
return log_probs
|
||||
|
||||
|
||||
def discretized_gaussian_log_likelihood(x, *, means, log_scales):
|
||||
"""
|
||||
Compute the log-likelihood of a Gaussian distribution discretizing to a
|
||||
given image.
|
||||
:param x: the target images. It is assumed that this was uint8 values,
|
||||
rescaled to the range [-1, 1].
|
||||
:param means: the Gaussian mean Tensor.
|
||||
:param log_scales: the Gaussian log stddev Tensor.
|
||||
:return: a tensor like x of log probabilities (in nats).
|
||||
"""
|
||||
assert x.shape == means.shape == log_scales.shape
|
||||
centered_x = x - means
|
||||
inv_stdv = th.exp(-log_scales)
|
||||
plus_in = inv_stdv * (centered_x + 1.0 / 255.0)
|
||||
cdf_plus = approx_standard_normal_cdf(plus_in)
|
||||
min_in = inv_stdv * (centered_x - 1.0 / 255.0)
|
||||
cdf_min = approx_standard_normal_cdf(min_in)
|
||||
log_cdf_plus = th.log(cdf_plus.clamp(min=1e-12))
|
||||
log_one_minus_cdf_min = th.log((1.0 - cdf_min).clamp(min=1e-12))
|
||||
cdf_delta = cdf_plus - cdf_min
|
||||
log_probs = th.where(
|
||||
x < -0.999,
|
||||
log_cdf_plus,
|
||||
th.where(x > 0.999, log_one_minus_cdf_min, th.log(cdf_delta.clamp(min=1e-12))),
|
||||
)
|
||||
assert log_probs.shape == x.shape
|
||||
return log_probs
|
||||
Executable
+829
@@ -0,0 +1,829 @@
|
||||
# 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
|
||||
|
||||
|
||||
import enum
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import torch as th
|
||||
|
||||
from .diffusion_utils import discretized_gaussian_log_likelihood, normal_kl
|
||||
|
||||
|
||||
def mean_flat(tensor):
|
||||
"""
|
||||
Take the mean over all non-batch dimensions.
|
||||
"""
|
||||
return tensor.mean(dim=list(range(1, len(tensor.shape))))
|
||||
|
||||
|
||||
class ModelMeanType(enum.Enum):
|
||||
"""
|
||||
Which type of output the model predicts.
|
||||
"""
|
||||
|
||||
PREVIOUS_X = enum.auto() # the model predicts x_{t-1}
|
||||
START_X = enum.auto() # the model predicts x_0
|
||||
EPSILON = enum.auto() # the model predicts epsilon
|
||||
|
||||
|
||||
class ModelVarType(enum.Enum):
|
||||
"""
|
||||
What is used as the model's output variance.
|
||||
The LEARNED_RANGE option has been added to allow the model to predict
|
||||
values between FIXED_SMALL and FIXED_LARGE, making its job easier.
|
||||
"""
|
||||
|
||||
LEARNED = enum.auto()
|
||||
FIXED_SMALL = enum.auto()
|
||||
FIXED_LARGE = enum.auto()
|
||||
LEARNED_RANGE = enum.auto()
|
||||
|
||||
|
||||
class LossType(enum.Enum):
|
||||
MSE = enum.auto() # use raw MSE loss (and KL when learning variances)
|
||||
RESCALED_MSE = enum.auto() # use raw MSE loss (with RESCALED_KL when learning variances)
|
||||
KL = enum.auto() # use the variational lower-bound
|
||||
RESCALED_KL = enum.auto() # like KL, but rescale to estimate the full VLB
|
||||
|
||||
def is_vb(self):
|
||||
return self == LossType.KL or self == LossType.RESCALED_KL
|
||||
|
||||
|
||||
def _warmup_beta(beta_start, beta_end, num_diffusion_timesteps, warmup_frac):
|
||||
betas = beta_end * np.ones(num_diffusion_timesteps, dtype=np.float64)
|
||||
warmup_time = int(num_diffusion_timesteps * warmup_frac)
|
||||
betas[:warmup_time] = np.linspace(beta_start, beta_end, warmup_time, dtype=np.float64)
|
||||
return betas
|
||||
|
||||
|
||||
def get_beta_schedule(beta_schedule, *, beta_start, beta_end, num_diffusion_timesteps):
|
||||
"""
|
||||
This is the deprecated API for creating beta schedules.
|
||||
See get_named_beta_schedule() for the new library of schedules.
|
||||
"""
|
||||
if beta_schedule == "quad":
|
||||
betas = (
|
||||
np.linspace(
|
||||
beta_start**0.5,
|
||||
beta_end**0.5,
|
||||
num_diffusion_timesteps,
|
||||
dtype=np.float64,
|
||||
)
|
||||
** 2
|
||||
)
|
||||
elif beta_schedule == "linear":
|
||||
betas = np.linspace(beta_start, beta_end, num_diffusion_timesteps, dtype=np.float64)
|
||||
elif beta_schedule == "warmup10":
|
||||
betas = _warmup_beta(beta_start, beta_end, num_diffusion_timesteps, 0.1)
|
||||
elif beta_schedule == "warmup50":
|
||||
betas = _warmup_beta(beta_start, beta_end, num_diffusion_timesteps, 0.5)
|
||||
elif beta_schedule == "const":
|
||||
betas = beta_end * np.ones(num_diffusion_timesteps, dtype=np.float64)
|
||||
elif beta_schedule == "jsd": # 1/T, 1/(T-1), 1/(T-2), ..., 1
|
||||
betas = 1.0 / np.linspace(num_diffusion_timesteps, 1, num_diffusion_timesteps, dtype=np.float64)
|
||||
else:
|
||||
raise NotImplementedError(beta_schedule)
|
||||
assert betas.shape == (num_diffusion_timesteps,)
|
||||
return betas
|
||||
|
||||
|
||||
def get_named_beta_schedule(schedule_name, num_diffusion_timesteps):
|
||||
"""
|
||||
Get a pre-defined beta schedule for the given name.
|
||||
The beta schedule library consists of beta schedules which remain similar
|
||||
in the limit of num_diffusion_timesteps.
|
||||
Beta schedules may be added, but should not be removed or changed once
|
||||
they are committed to maintain backwards compatibility.
|
||||
"""
|
||||
if schedule_name == "linear":
|
||||
# Linear schedule from Ho et al, extended to work for any number of
|
||||
# diffusion steps.
|
||||
scale = 1000 / num_diffusion_timesteps
|
||||
return get_beta_schedule(
|
||||
"linear",
|
||||
beta_start=scale * 0.0001,
|
||||
beta_end=scale * 0.02,
|
||||
num_diffusion_timesteps=num_diffusion_timesteps,
|
||||
)
|
||||
elif schedule_name == "squaredcos_cap_v2":
|
||||
return betas_for_alpha_bar(
|
||||
num_diffusion_timesteps,
|
||||
lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"unknown beta schedule: {schedule_name}")
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
class GaussianDiffusion:
|
||||
"""
|
||||
Utilities for training and sampling diffusion models.
|
||||
Original ported from this codebase:
|
||||
https://github.com/hojonathanho/diffusion/blob/1e0dceb3b3495bbe19116a5e1b3596cd0706c543/diffusion_tf/diffusion_utils_2.py#L42
|
||||
:param betas: a 1-D numpy array of betas for each diffusion timestep,
|
||||
starting at T and going to 1.
|
||||
"""
|
||||
|
||||
def __init__(self, *, betas, model_mean_type, model_var_type, loss_type):
|
||||
self.model_mean_type = model_mean_type
|
||||
self.model_var_type = model_var_type
|
||||
self.loss_type = loss_type
|
||||
|
||||
# Use float64 for accuracy.
|
||||
betas = np.array(betas, dtype=np.float64)
|
||||
self.betas = betas
|
||||
assert len(betas.shape) == 1, "betas must be 1-D"
|
||||
assert (betas > 0).all() and (betas <= 1).all()
|
||||
|
||||
self.num_timesteps = int(betas.shape[0])
|
||||
|
||||
alphas = 1.0 - betas
|
||||
self.alphas_cumprod = np.cumprod(alphas, axis=0)
|
||||
self.alphas_cumprod_prev = np.append(1.0, self.alphas_cumprod[:-1])
|
||||
self.alphas_cumprod_next = np.append(self.alphas_cumprod[1:], 0.0)
|
||||
assert self.alphas_cumprod_prev.shape == (self.num_timesteps,)
|
||||
|
||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
||||
self.sqrt_alphas_cumprod = np.sqrt(self.alphas_cumprod)
|
||||
self.sqrt_one_minus_alphas_cumprod = np.sqrt(1.0 - self.alphas_cumprod)
|
||||
self.log_one_minus_alphas_cumprod = np.log(1.0 - self.alphas_cumprod)
|
||||
self.sqrt_recip_alphas_cumprod = np.sqrt(1.0 / self.alphas_cumprod)
|
||||
self.sqrt_recipm1_alphas_cumprod = np.sqrt(1.0 / self.alphas_cumprod - 1)
|
||||
|
||||
# calculations for posterior q(x_{t-1} | x_t, x_0)
|
||||
self.posterior_variance = betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
|
||||
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
|
||||
self.posterior_log_variance_clipped = (
|
||||
np.log(np.append(self.posterior_variance[1], self.posterior_variance[1:]))
|
||||
if len(self.posterior_variance) > 1
|
||||
else np.array([])
|
||||
)
|
||||
|
||||
self.posterior_mean_coef1 = betas * np.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
|
||||
self.posterior_mean_coef2 = (1.0 - self.alphas_cumprod_prev) * np.sqrt(alphas) / (1.0 - self.alphas_cumprod)
|
||||
|
||||
def q_mean_variance(self, x_start, t):
|
||||
"""
|
||||
Get the distribution q(x_t | x_0).
|
||||
:param x_start: the [N x C x ...] tensor of noiseless inputs.
|
||||
:param t: the number of diffusion steps (minus 1). Here, 0 means one step.
|
||||
:return: A tuple (mean, variance, log_variance), all of x_start's shape.
|
||||
"""
|
||||
mean = _extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
|
||||
variance = _extract_into_tensor(1.0 - self.alphas_cumprod, t, x_start.shape)
|
||||
log_variance = _extract_into_tensor(self.log_one_minus_alphas_cumprod, t, x_start.shape)
|
||||
return mean, variance, log_variance
|
||||
|
||||
def q_sample(self, x_start, t, noise=None):
|
||||
"""
|
||||
Diffuse the data for a given number of diffusion steps.
|
||||
In other words, sample from q(x_t | x_0).
|
||||
:param x_start: the initial data batch.
|
||||
:param t: the number of diffusion steps (minus 1). Here, 0 means one step.
|
||||
:param noise: if specified, the split-out normal noise.
|
||||
:return: A noisy version of x_start.
|
||||
"""
|
||||
if noise is None:
|
||||
noise = th.randn_like(x_start)
|
||||
assert noise.shape == x_start.shape
|
||||
return (
|
||||
_extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
|
||||
+ _extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
|
||||
)
|
||||
|
||||
def q_posterior_mean_variance(self, x_start, x_t, t):
|
||||
"""
|
||||
Compute the mean and variance of the diffusion posterior:
|
||||
q(x_{t-1} | x_t, x_0)
|
||||
"""
|
||||
assert x_start.shape == x_t.shape
|
||||
posterior_mean = (
|
||||
_extract_into_tensor(self.posterior_mean_coef1, t, x_t.shape) * x_start
|
||||
+ _extract_into_tensor(self.posterior_mean_coef2, t, x_t.shape) * x_t
|
||||
)
|
||||
posterior_variance = _extract_into_tensor(self.posterior_variance, t, x_t.shape)
|
||||
posterior_log_variance_clipped = _extract_into_tensor(self.posterior_log_variance_clipped, t, x_t.shape)
|
||||
assert (
|
||||
posterior_mean.shape[0]
|
||||
== posterior_variance.shape[0]
|
||||
== posterior_log_variance_clipped.shape[0]
|
||||
== x_start.shape[0]
|
||||
)
|
||||
return posterior_mean, posterior_variance, posterior_log_variance_clipped
|
||||
|
||||
def p_mean_variance(self, model, x, t, clip_denoised=True, denoised_fn=None, model_kwargs=None):
|
||||
"""
|
||||
Apply the model to get p(x_{t-1} | x_t), as well as a prediction of
|
||||
the initial x, x_0.
|
||||
:param model: the model, which takes a signal and a batch of timesteps
|
||||
as input.
|
||||
:param x: the [N x C x ...] tensor at time t.
|
||||
:param t: a 1-D Tensor of timesteps.
|
||||
:param clip_denoised: if True, clip the denoised signal into [-1, 1].
|
||||
:param denoised_fn: if not None, a function which applies to the
|
||||
x_start prediction before it is used to sample. Applies before
|
||||
clip_denoised.
|
||||
:param model_kwargs: if not None, a dict of extra keyword arguments to
|
||||
pass to the model. This can be used for conditioning.
|
||||
:return: a dict with the following keys:
|
||||
- 'mean': the model mean output.
|
||||
- 'variance': the model variance output.
|
||||
- 'log_variance': the log of 'variance'.
|
||||
- 'pred_xstart': the prediction for x_0.
|
||||
"""
|
||||
if model_kwargs is None:
|
||||
model_kwargs = {}
|
||||
|
||||
B, C = x.shape[:2]
|
||||
assert t.shape == (B,)
|
||||
model_output = model(x, t, **model_kwargs)
|
||||
if isinstance(model_output, tuple):
|
||||
model_output, extra = model_output
|
||||
else:
|
||||
extra = None
|
||||
|
||||
if self.model_var_type in [ModelVarType.LEARNED, ModelVarType.LEARNED_RANGE]:
|
||||
assert model_output.shape == (B, C * 2, *x.shape[2:])
|
||||
model_output, model_var_values = th.split(model_output, C, dim=1)
|
||||
min_log = _extract_into_tensor(self.posterior_log_variance_clipped, t, x.shape)
|
||||
max_log = _extract_into_tensor(np.log(self.betas), t, x.shape)
|
||||
# The model_var_values is [-1, 1] for [min_var, max_var].
|
||||
frac = (model_var_values + 1) / 2
|
||||
model_log_variance = frac * max_log + (1 - frac) * min_log
|
||||
model_variance = th.exp(model_log_variance)
|
||||
else:
|
||||
model_variance, model_log_variance = {
|
||||
# for fixedlarge, we set the initial (log-)variance like so
|
||||
# to get a better decoder log likelihood.
|
||||
ModelVarType.FIXED_LARGE: (
|
||||
np.append(self.posterior_variance[1], self.betas[1:]),
|
||||
np.log(np.append(self.posterior_variance[1], self.betas[1:])),
|
||||
),
|
||||
ModelVarType.FIXED_SMALL: (
|
||||
self.posterior_variance,
|
||||
self.posterior_log_variance_clipped,
|
||||
),
|
||||
}[self.model_var_type]
|
||||
model_variance = _extract_into_tensor(model_variance, t, x.shape)
|
||||
model_log_variance = _extract_into_tensor(model_log_variance, t, x.shape)
|
||||
|
||||
def process_xstart(x):
|
||||
if denoised_fn is not None:
|
||||
x = denoised_fn(x)
|
||||
if clip_denoised:
|
||||
return x.clamp(-1, 1)
|
||||
return x
|
||||
|
||||
if self.model_mean_type == ModelMeanType.START_X:
|
||||
pred_xstart = process_xstart(model_output)
|
||||
else:
|
||||
pred_xstart = process_xstart(self._predict_xstart_from_eps(x_t=x, t=t, eps=model_output))
|
||||
model_mean, _, _ = self.q_posterior_mean_variance(x_start=pred_xstart, x_t=x, t=t)
|
||||
|
||||
assert model_mean.shape == model_log_variance.shape == pred_xstart.shape == x.shape
|
||||
return {
|
||||
"mean": model_mean,
|
||||
"variance": model_variance,
|
||||
"log_variance": model_log_variance,
|
||||
"pred_xstart": pred_xstart,
|
||||
"extra": extra,
|
||||
}
|
||||
|
||||
def _predict_xstart_from_eps(self, x_t, t, eps):
|
||||
assert x_t.shape == eps.shape
|
||||
return (
|
||||
_extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t
|
||||
- _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * eps
|
||||
)
|
||||
|
||||
def _predict_eps_from_xstart(self, x_t, t, pred_xstart):
|
||||
return (
|
||||
_extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - pred_xstart
|
||||
) / _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
|
||||
|
||||
def condition_mean(self, cond_fn, p_mean_var, x, t, model_kwargs=None):
|
||||
"""
|
||||
Compute the mean for the previous step, given a function cond_fn that
|
||||
computes the gradient of a conditional log probability with respect to
|
||||
x. In particular, cond_fn computes grad(log(p(y|x))), and we want to
|
||||
condition on y.
|
||||
This uses the conditioning strategy from Sohl-Dickstein et al. (2015).
|
||||
"""
|
||||
gradient = cond_fn(x, t, **model_kwargs)
|
||||
new_mean = p_mean_var["mean"].float() + p_mean_var["variance"] * gradient.float()
|
||||
return new_mean
|
||||
|
||||
def condition_score(self, cond_fn, p_mean_var, x, t, model_kwargs=None):
|
||||
"""
|
||||
Compute what the p_mean_variance output would have been, should the
|
||||
model's score function be conditioned by cond_fn.
|
||||
See condition_mean() for details on cond_fn.
|
||||
Unlike condition_mean(), this instead uses the conditioning strategy
|
||||
from Song et al (2020).
|
||||
"""
|
||||
alpha_bar = _extract_into_tensor(self.alphas_cumprod, t, x.shape)
|
||||
|
||||
eps = self._predict_eps_from_xstart(x, t, p_mean_var["pred_xstart"])
|
||||
eps = eps - (1 - alpha_bar).sqrt() * cond_fn(x, t, **model_kwargs)
|
||||
|
||||
out = p_mean_var.copy()
|
||||
out["pred_xstart"] = self._predict_xstart_from_eps(x, t, eps)
|
||||
out["mean"], _, _ = self.q_posterior_mean_variance(x_start=out["pred_xstart"], x_t=x, t=t)
|
||||
return out
|
||||
|
||||
def p_sample(
|
||||
self,
|
||||
model,
|
||||
x,
|
||||
t,
|
||||
clip_denoised=True,
|
||||
denoised_fn=None,
|
||||
cond_fn=None,
|
||||
model_kwargs=None,
|
||||
):
|
||||
"""
|
||||
Sample x_{t-1} from the model at the given timestep.
|
||||
:param model: the model to sample from.
|
||||
:param x: the current tensor at x_{t-1}.
|
||||
:param t: the value of t, starting at 0 for the first diffusion step.
|
||||
:param clip_denoised: if True, clip the x_start prediction to [-1, 1].
|
||||
:param denoised_fn: if not None, a function which applies to the
|
||||
x_start prediction before it is used to sample.
|
||||
:param cond_fn: if not None, this is a gradient function that acts
|
||||
similarly to the model.
|
||||
:param model_kwargs: if not None, a dict of extra keyword arguments to
|
||||
pass to the model. This can be used for conditioning.
|
||||
:return: a dict containing the following keys:
|
||||
- 'sample': a random sample from the model.
|
||||
- 'pred_xstart': a prediction of x_0.
|
||||
"""
|
||||
out = self.p_mean_variance(
|
||||
model,
|
||||
x,
|
||||
t,
|
||||
clip_denoised=clip_denoised,
|
||||
denoised_fn=denoised_fn,
|
||||
model_kwargs=model_kwargs,
|
||||
)
|
||||
noise = th.randn_like(x)
|
||||
nonzero_mask = (t != 0).float().view(-1, *([1] * (len(x.shape) - 1))) # no noise when t == 0
|
||||
if cond_fn is not None:
|
||||
out["mean"] = self.condition_mean(cond_fn, out, x, t, model_kwargs=model_kwargs)
|
||||
sample = out["mean"] + nonzero_mask * th.exp(0.5 * out["log_variance"]) * noise
|
||||
return {"sample": sample, "pred_xstart": out["pred_xstart"]}
|
||||
|
||||
def p_sample_loop(
|
||||
self,
|
||||
model,
|
||||
shape,
|
||||
noise=None,
|
||||
clip_denoised=True,
|
||||
denoised_fn=None,
|
||||
cond_fn=None,
|
||||
model_kwargs=None,
|
||||
device=None,
|
||||
progress=False,
|
||||
):
|
||||
"""
|
||||
Generate samples from the model.
|
||||
:param model: the model module.
|
||||
:param shape: the shape of the samples, (N, C, H, W).
|
||||
:param noise: if specified, the noise from the encoder to sample.
|
||||
Should be of the same shape as `shape`.
|
||||
:param clip_denoised: if True, clip x_start predictions to [-1, 1].
|
||||
:param denoised_fn: if not None, a function which applies to the
|
||||
x_start prediction before it is used to sample.
|
||||
:param cond_fn: if not None, this is a gradient function that acts
|
||||
similarly to the model.
|
||||
:param model_kwargs: if not None, a dict of extra keyword arguments to
|
||||
pass to the model. This can be used for conditioning.
|
||||
:param device: if specified, the device to create the samples on.
|
||||
If not specified, use a model parameter's device.
|
||||
:param progress: if True, show a tqdm progress bar.
|
||||
:return: a non-differentiable batch of samples.
|
||||
"""
|
||||
final = None
|
||||
for sample in self.p_sample_loop_progressive(
|
||||
model,
|
||||
shape,
|
||||
noise=noise,
|
||||
clip_denoised=clip_denoised,
|
||||
denoised_fn=denoised_fn,
|
||||
cond_fn=cond_fn,
|
||||
model_kwargs=model_kwargs,
|
||||
device=device,
|
||||
progress=progress,
|
||||
):
|
||||
final = sample
|
||||
return final["sample"]
|
||||
|
||||
def p_sample_loop_progressive(
|
||||
self,
|
||||
model,
|
||||
shape,
|
||||
noise=None,
|
||||
clip_denoised=True,
|
||||
denoised_fn=None,
|
||||
cond_fn=None,
|
||||
model_kwargs=None,
|
||||
device=None,
|
||||
progress=False,
|
||||
):
|
||||
"""
|
||||
Generate samples from the model and yield intermediate samples from
|
||||
each timestep of diffusion.
|
||||
Arguments are the same as p_sample_loop().
|
||||
Returns a generator over dicts, where each dict is the return value of
|
||||
p_sample().
|
||||
"""
|
||||
if device is None:
|
||||
device = next(model.parameters()).device
|
||||
assert isinstance(shape, (tuple, list))
|
||||
if noise is not None:
|
||||
img = noise
|
||||
else:
|
||||
img = th.randn(*shape, device=device)
|
||||
indices = list(range(self.num_timesteps))[::-1]
|
||||
|
||||
if progress:
|
||||
# Lazy import so that we don't depend on tqdm.
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
indices = tqdm(indices)
|
||||
|
||||
for i in indices:
|
||||
t = th.tensor([i] * shape[0], device=device)
|
||||
with th.no_grad():
|
||||
out = self.p_sample(
|
||||
model,
|
||||
img,
|
||||
t,
|
||||
clip_denoised=clip_denoised,
|
||||
denoised_fn=denoised_fn,
|
||||
cond_fn=cond_fn,
|
||||
model_kwargs=model_kwargs,
|
||||
)
|
||||
yield out
|
||||
img = out["sample"]
|
||||
|
||||
def ddim_sample(
|
||||
self,
|
||||
model,
|
||||
x,
|
||||
t,
|
||||
clip_denoised=True,
|
||||
denoised_fn=None,
|
||||
cond_fn=None,
|
||||
model_kwargs=None,
|
||||
eta=0.0,
|
||||
):
|
||||
"""
|
||||
Sample x_{t-1} from the model using DDIM.
|
||||
Same usage as p_sample().
|
||||
"""
|
||||
out = self.p_mean_variance(
|
||||
model,
|
||||
x,
|
||||
t,
|
||||
clip_denoised=clip_denoised,
|
||||
denoised_fn=denoised_fn,
|
||||
model_kwargs=model_kwargs,
|
||||
)
|
||||
if cond_fn is not None:
|
||||
out = self.condition_score(cond_fn, out, x, t, model_kwargs=model_kwargs)
|
||||
|
||||
# Usually our model outputs epsilon, but we re-derive it
|
||||
# in case we used x_start or x_prev prediction.
|
||||
eps = self._predict_eps_from_xstart(x, t, out["pred_xstart"])
|
||||
|
||||
alpha_bar = _extract_into_tensor(self.alphas_cumprod, t, x.shape)
|
||||
alpha_bar_prev = _extract_into_tensor(self.alphas_cumprod_prev, t, x.shape)
|
||||
sigma = eta * th.sqrt((1 - alpha_bar_prev) / (1 - alpha_bar)) * th.sqrt(1 - alpha_bar / alpha_bar_prev)
|
||||
# Equation 12.
|
||||
noise = th.randn_like(x)
|
||||
mean_pred = out["pred_xstart"] * th.sqrt(alpha_bar_prev) + th.sqrt(1 - alpha_bar_prev - sigma**2) * eps
|
||||
nonzero_mask = (t != 0).float().view(-1, *([1] * (len(x.shape) - 1))) # no noise when t == 0
|
||||
sample = mean_pred + nonzero_mask * sigma * noise
|
||||
return {"sample": sample, "pred_xstart": out["pred_xstart"]}
|
||||
|
||||
def ddim_reverse_sample(
|
||||
self,
|
||||
model,
|
||||
x,
|
||||
t,
|
||||
clip_denoised=True,
|
||||
denoised_fn=None,
|
||||
cond_fn=None,
|
||||
model_kwargs=None,
|
||||
eta=0.0,
|
||||
):
|
||||
"""
|
||||
Sample x_{t+1} from the model using DDIM reverse ODE.
|
||||
"""
|
||||
assert eta == 0.0, "Reverse ODE only for deterministic path"
|
||||
out = self.p_mean_variance(
|
||||
model,
|
||||
x,
|
||||
t,
|
||||
clip_denoised=clip_denoised,
|
||||
denoised_fn=denoised_fn,
|
||||
model_kwargs=model_kwargs,
|
||||
)
|
||||
if cond_fn is not None:
|
||||
out = self.condition_score(cond_fn, out, x, t, model_kwargs=model_kwargs)
|
||||
# Usually our model outputs epsilon, but we re-derive it
|
||||
# in case we used x_start or x_prev prediction.
|
||||
eps = (
|
||||
_extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x.shape) * x - out["pred_xstart"]
|
||||
) / _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x.shape)
|
||||
alpha_bar_next = _extract_into_tensor(self.alphas_cumprod_next, t, x.shape)
|
||||
|
||||
# Equation 12. reversed
|
||||
mean_pred = out["pred_xstart"] * th.sqrt(alpha_bar_next) + th.sqrt(1 - alpha_bar_next) * eps
|
||||
|
||||
return {"sample": mean_pred, "pred_xstart": out["pred_xstart"]}
|
||||
|
||||
def ddim_sample_loop(
|
||||
self,
|
||||
model,
|
||||
shape,
|
||||
noise=None,
|
||||
clip_denoised=True,
|
||||
denoised_fn=None,
|
||||
cond_fn=None,
|
||||
model_kwargs=None,
|
||||
device=None,
|
||||
progress=False,
|
||||
eta=0.0,
|
||||
):
|
||||
"""
|
||||
Generate samples from the model using DDIM.
|
||||
Same usage as p_sample_loop().
|
||||
"""
|
||||
final = None
|
||||
for sample in self.ddim_sample_loop_progressive(
|
||||
model,
|
||||
shape,
|
||||
noise=noise,
|
||||
clip_denoised=clip_denoised,
|
||||
denoised_fn=denoised_fn,
|
||||
cond_fn=cond_fn,
|
||||
model_kwargs=model_kwargs,
|
||||
device=device,
|
||||
progress=progress,
|
||||
eta=eta,
|
||||
):
|
||||
final = sample
|
||||
return final["sample"]
|
||||
|
||||
def ddim_sample_loop_progressive(
|
||||
self,
|
||||
model,
|
||||
shape,
|
||||
noise=None,
|
||||
clip_denoised=True,
|
||||
denoised_fn=None,
|
||||
cond_fn=None,
|
||||
model_kwargs=None,
|
||||
device=None,
|
||||
progress=False,
|
||||
eta=0.0,
|
||||
):
|
||||
"""
|
||||
Use DDIM to sample from the model and yield intermediate samples from
|
||||
each timestep of DDIM.
|
||||
Same usage as p_sample_loop_progressive().
|
||||
"""
|
||||
if device is None:
|
||||
device = next(model.parameters()).device
|
||||
assert isinstance(shape, (tuple, list))
|
||||
if noise is not None:
|
||||
img = noise
|
||||
else:
|
||||
img = th.randn(*shape, device=device)
|
||||
indices = list(range(self.num_timesteps))[::-1]
|
||||
|
||||
if progress:
|
||||
# Lazy import so that we don't depend on tqdm.
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
indices = tqdm(indices)
|
||||
|
||||
for i in indices:
|
||||
t = th.tensor([i] * shape[0], device=device)
|
||||
with th.no_grad():
|
||||
out = self.ddim_sample(
|
||||
model,
|
||||
img,
|
||||
t,
|
||||
clip_denoised=clip_denoised,
|
||||
denoised_fn=denoised_fn,
|
||||
cond_fn=cond_fn,
|
||||
model_kwargs=model_kwargs,
|
||||
eta=eta,
|
||||
)
|
||||
yield out
|
||||
img = out["sample"]
|
||||
|
||||
def _vb_terms_bpd(self, model, x_start, x_t, t, clip_denoised=True, model_kwargs=None):
|
||||
"""
|
||||
Get a term for the variational lower-bound.
|
||||
The resulting units are bits (rather than nats, as one might expect).
|
||||
This allows for comparison to other papers.
|
||||
:return: a dict with the following keys:
|
||||
- 'output': a shape [N] tensor of NLLs or KLs.
|
||||
- 'pred_xstart': the x_0 predictions.
|
||||
"""
|
||||
true_mean, _, true_log_variance_clipped = self.q_posterior_mean_variance(x_start=x_start, x_t=x_t, t=t)
|
||||
out = self.p_mean_variance(model, x_t, t, clip_denoised=clip_denoised, model_kwargs=model_kwargs)
|
||||
kl = normal_kl(true_mean, true_log_variance_clipped, out["mean"], out["log_variance"])
|
||||
kl = mean_flat(kl) / np.log(2.0)
|
||||
|
||||
decoder_nll = -discretized_gaussian_log_likelihood(
|
||||
x_start, means=out["mean"], log_scales=0.5 * out["log_variance"]
|
||||
)
|
||||
assert decoder_nll.shape == x_start.shape
|
||||
decoder_nll = mean_flat(decoder_nll) / np.log(2.0)
|
||||
|
||||
# At the first timestep return the decoder NLL,
|
||||
# otherwise return KL(q(x_{t-1}|x_t,x_0) || p(x_{t-1}|x_t))
|
||||
output = th.where((t == 0), decoder_nll, kl)
|
||||
return {"output": output, "pred_xstart": out["pred_xstart"]}
|
||||
|
||||
def training_losses(self, model, x_start, t, model_kwargs=None, noise=None):
|
||||
"""
|
||||
Compute training losses for a single timestep.
|
||||
:param model: the model to evaluate loss on.
|
||||
:param x_start: the [N x C x ...] tensor of inputs.
|
||||
:param t: a batch of timestep indices.
|
||||
:param model_kwargs: if not None, a dict of extra keyword arguments to
|
||||
pass to the model. This can be used for conditioning.
|
||||
:param noise: if specified, the specific Gaussian noise to try to remove.
|
||||
:return: a dict with the key "loss" containing a tensor of shape [N].
|
||||
Some mean or variance settings may also have other keys.
|
||||
"""
|
||||
if model_kwargs is None:
|
||||
model_kwargs = {}
|
||||
if noise is None:
|
||||
noise = th.randn_like(x_start)
|
||||
x_t = self.q_sample(x_start, t, noise=noise)
|
||||
|
||||
terms = {}
|
||||
|
||||
if self.loss_type == LossType.KL or self.loss_type == LossType.RESCALED_KL:
|
||||
terms["loss"] = self._vb_terms_bpd(
|
||||
model=model,
|
||||
x_start=x_start,
|
||||
x_t=x_t,
|
||||
t=t,
|
||||
clip_denoised=False,
|
||||
model_kwargs=model_kwargs,
|
||||
)["output"]
|
||||
if self.loss_type == LossType.RESCALED_KL:
|
||||
terms["loss"] *= self.num_timesteps
|
||||
elif self.loss_type == LossType.MSE or self.loss_type == LossType.RESCALED_MSE:
|
||||
model_output = model(x_t, t, **model_kwargs)
|
||||
|
||||
if self.model_var_type in [
|
||||
ModelVarType.LEARNED,
|
||||
ModelVarType.LEARNED_RANGE,
|
||||
]:
|
||||
B, C = x_t.shape[:2]
|
||||
assert model_output.shape == (B, C * 2, *x_t.shape[2:])
|
||||
model_output, model_var_values = th.split(model_output, C, dim=1)
|
||||
# Learn the variance using the variational bound, but don't let
|
||||
# it affect our mean prediction.
|
||||
frozen_out = th.cat([model_output.detach(), model_var_values], dim=1)
|
||||
terms["vb"] = self._vb_terms_bpd(
|
||||
model=lambda *args, r=frozen_out: r,
|
||||
x_start=x_start,
|
||||
x_t=x_t,
|
||||
t=t,
|
||||
clip_denoised=False,
|
||||
)["output"]
|
||||
if self.loss_type == LossType.RESCALED_MSE:
|
||||
# Divide by 1000 for equivalence with initial implementation.
|
||||
# Without a factor of 1/1000, the VB term hurts the MSE term.
|
||||
terms["vb"] *= self.num_timesteps / 1000.0
|
||||
|
||||
target = {
|
||||
ModelMeanType.PREVIOUS_X: self.q_posterior_mean_variance(x_start=x_start, x_t=x_t, t=t)[0],
|
||||
ModelMeanType.START_X: x_start,
|
||||
ModelMeanType.EPSILON: noise,
|
||||
}[self.model_mean_type]
|
||||
assert model_output.shape == target.shape == x_start.shape
|
||||
terms["mse"] = mean_flat((target - model_output) ** 2)
|
||||
if "vb" in terms:
|
||||
terms["loss"] = terms["mse"] + terms["vb"]
|
||||
else:
|
||||
terms["loss"] = terms["mse"]
|
||||
else:
|
||||
raise NotImplementedError(self.loss_type)
|
||||
|
||||
return terms
|
||||
|
||||
def _prior_bpd(self, x_start):
|
||||
"""
|
||||
Get the prior KL term for the variational lower-bound, measured in
|
||||
bits-per-dim.
|
||||
This term can't be optimized, as it only depends on the encoder.
|
||||
:param x_start: the [N x C x ...] tensor of inputs.
|
||||
:return: a batch of [N] KL values (in bits), one per batch element.
|
||||
"""
|
||||
batch_size = x_start.shape[0]
|
||||
t = th.tensor([self.num_timesteps - 1] * batch_size, device=x_start.device)
|
||||
qt_mean, _, qt_log_variance = self.q_mean_variance(x_start, t)
|
||||
kl_prior = normal_kl(mean1=qt_mean, logvar1=qt_log_variance, mean2=0.0, logvar2=0.0)
|
||||
return mean_flat(kl_prior) / np.log(2.0)
|
||||
|
||||
def calc_bpd_loop(self, model, x_start, clip_denoised=True, model_kwargs=None):
|
||||
"""
|
||||
Compute the entire variational lower-bound, measured in bits-per-dim,
|
||||
as well as other related quantities.
|
||||
:param model: the model to evaluate loss on.
|
||||
:param x_start: the [N x C x ...] tensor of inputs.
|
||||
:param clip_denoised: if True, clip denoised samples.
|
||||
:param model_kwargs: if not None, a dict of extra keyword arguments to
|
||||
pass to the model. This can be used for conditioning.
|
||||
:return: a dict containing the following keys:
|
||||
- total_bpd: the total variational lower-bound, per batch element.
|
||||
- prior_bpd: the prior term in the lower-bound.
|
||||
- vb: an [N x T] tensor of terms in the lower-bound.
|
||||
- xstart_mse: an [N x T] tensor of x_0 MSEs for each timestep.
|
||||
- mse: an [N x T] tensor of epsilon MSEs for each timestep.
|
||||
"""
|
||||
device = x_start.device
|
||||
batch_size = x_start.shape[0]
|
||||
|
||||
vb = []
|
||||
xstart_mse = []
|
||||
mse = []
|
||||
for t in list(range(self.num_timesteps))[::-1]:
|
||||
t_batch = th.tensor([t] * batch_size, device=device)
|
||||
noise = th.randn_like(x_start)
|
||||
x_t = self.q_sample(x_start=x_start, t=t_batch, noise=noise)
|
||||
# Calculate VLB term at the current timestep
|
||||
with th.no_grad():
|
||||
out = self._vb_terms_bpd(
|
||||
model,
|
||||
x_start=x_start,
|
||||
x_t=x_t,
|
||||
t=t_batch,
|
||||
clip_denoised=clip_denoised,
|
||||
model_kwargs=model_kwargs,
|
||||
)
|
||||
vb.append(out["output"])
|
||||
xstart_mse.append(mean_flat((out["pred_xstart"] - x_start) ** 2))
|
||||
eps = self._predict_eps_from_xstart(x_t, t_batch, out["pred_xstart"])
|
||||
mse.append(mean_flat((eps - noise) ** 2))
|
||||
|
||||
vb = th.stack(vb, dim=1)
|
||||
xstart_mse = th.stack(xstart_mse, dim=1)
|
||||
mse = th.stack(mse, dim=1)
|
||||
|
||||
prior_bpd = self._prior_bpd(x_start)
|
||||
total_bpd = vb.sum(dim=1) + prior_bpd
|
||||
return {
|
||||
"total_bpd": total_bpd,
|
||||
"prior_bpd": prior_bpd,
|
||||
"vb": vb,
|
||||
"xstart_mse": xstart_mse,
|
||||
"mse": mse,
|
||||
}
|
||||
|
||||
|
||||
def _extract_into_tensor(arr, timesteps, broadcast_shape):
|
||||
"""
|
||||
Extract values from a 1-D numpy array for a batch of indices.
|
||||
:param arr: the 1-D numpy array.
|
||||
:param timesteps: a tensor of indices into the array to extract.
|
||||
:param broadcast_shape: a larger shape of K dimensions with the batch
|
||||
dimension equal to the length of timesteps.
|
||||
:return: a tensor of shape [batch_size, 1, ...] where the shape has K dims.
|
||||
"""
|
||||
res = th.from_numpy(arr).to(device=timesteps.device)[timesteps].float()
|
||||
while len(res.shape) < len(broadcast_shape):
|
||||
res = res[..., None]
|
||||
return res + th.zeros(broadcast_shape, device=timesteps.device)
|
||||
Executable
+119
@@ -0,0 +1,119 @@
|
||||
# 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
|
||||
|
||||
import numpy as np
|
||||
import torch as th
|
||||
|
||||
from .gaussian_diffusion import GaussianDiffusion
|
||||
|
||||
|
||||
def space_timesteps(num_timesteps, section_counts):
|
||||
"""
|
||||
Create a list of timesteps to use from an original diffusion process,
|
||||
given the number of timesteps we want to take from equally-sized portions
|
||||
of the original process.
|
||||
For example, if there's 300 timesteps and the section counts are [10,15,20]
|
||||
then the first 100 timesteps are strided to be 10 timesteps, the second 100
|
||||
are strided to be 15 timesteps, and the final 100 are strided to be 20.
|
||||
If the stride is a string starting with "ddim", then the fixed striding
|
||||
from the DDIM paper is used, and only one section is allowed.
|
||||
:param num_timesteps: the number of diffusion steps in the original
|
||||
process to divide up.
|
||||
:param section_counts: either a list of numbers, or a string containing
|
||||
comma-separated numbers, indicating the step count
|
||||
per section. As a special case, use "ddimN" where N
|
||||
is a number of steps to use the striding from the
|
||||
DDIM paper.
|
||||
:return: a set of diffusion steps from the original process to use.
|
||||
"""
|
||||
if isinstance(section_counts, str):
|
||||
if section_counts.startswith("ddim"):
|
||||
desired_count = int(section_counts[len("ddim") :])
|
||||
for i in range(1, num_timesteps):
|
||||
if len(range(0, num_timesteps, i)) == desired_count:
|
||||
return set(range(0, num_timesteps, i))
|
||||
raise ValueError(f"cannot create exactly {num_timesteps} steps with an integer stride")
|
||||
section_counts = [int(x) for x in section_counts.split(",")]
|
||||
size_per = num_timesteps // len(section_counts)
|
||||
extra = num_timesteps % len(section_counts)
|
||||
start_idx = 0
|
||||
all_steps = []
|
||||
for i, section_count in enumerate(section_counts):
|
||||
size = size_per + (1 if i < extra else 0)
|
||||
if size < section_count:
|
||||
raise ValueError(f"cannot divide section of {size} steps into {section_count}")
|
||||
if section_count <= 1:
|
||||
frac_stride = 1
|
||||
else:
|
||||
frac_stride = (size - 1) / (section_count - 1)
|
||||
cur_idx = 0.0
|
||||
taken_steps = []
|
||||
for _ in range(section_count):
|
||||
taken_steps.append(start_idx + round(cur_idx))
|
||||
cur_idx += frac_stride
|
||||
all_steps += taken_steps
|
||||
start_idx += size
|
||||
return set(all_steps)
|
||||
|
||||
|
||||
class SpacedDiffusion(GaussianDiffusion):
|
||||
"""
|
||||
A diffusion process which can skip steps in a base diffusion process.
|
||||
:param use_timesteps: a collection (sequence or set) of timesteps from the
|
||||
original diffusion process to retain.
|
||||
:param kwargs: the kwargs to create the base diffusion process.
|
||||
"""
|
||||
|
||||
def __init__(self, use_timesteps, **kwargs):
|
||||
self.use_timesteps = set(use_timesteps)
|
||||
self.timestep_map = []
|
||||
self.original_num_steps = len(kwargs["betas"])
|
||||
|
||||
base_diffusion = GaussianDiffusion(**kwargs) # pylint: disable=missing-kwoa
|
||||
last_alpha_cumprod = 1.0
|
||||
new_betas = []
|
||||
for i, alpha_cumprod in enumerate(base_diffusion.alphas_cumprod):
|
||||
if i in self.use_timesteps:
|
||||
new_betas.append(1 - alpha_cumprod / last_alpha_cumprod)
|
||||
last_alpha_cumprod = alpha_cumprod
|
||||
self.timestep_map.append(i)
|
||||
kwargs["betas"] = np.array(new_betas)
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def p_mean_variance(self, model, *args, **kwargs): # pylint: disable=signature-differs
|
||||
return super().p_mean_variance(self._wrap_model(model), *args, **kwargs)
|
||||
|
||||
def training_losses(self, model, *args, **kwargs): # pylint: disable=signature-differs
|
||||
return super().training_losses(self._wrap_model(model), *args, **kwargs)
|
||||
|
||||
def condition_mean(self, cond_fn, *args, **kwargs):
|
||||
return super().condition_mean(self._wrap_model(cond_fn), *args, **kwargs)
|
||||
|
||||
def condition_score(self, cond_fn, *args, **kwargs):
|
||||
return super().condition_score(self._wrap_model(cond_fn), *args, **kwargs)
|
||||
|
||||
def _wrap_model(self, model):
|
||||
if isinstance(model, _WrappedModel):
|
||||
return model
|
||||
return _WrappedModel(model, self.timestep_map, self.original_num_steps)
|
||||
|
||||
def _scale_timesteps(self, t):
|
||||
# Scaling is done by the wrapped model.
|
||||
return t
|
||||
|
||||
|
||||
class _WrappedModel:
|
||||
def __init__(self, model, timestep_map, original_num_steps):
|
||||
self.model = model
|
||||
self.timestep_map = timestep_map
|
||||
# self.rescale_timesteps = rescale_timesteps
|
||||
self.original_num_steps = original_num_steps
|
||||
|
||||
def __call__(self, x, ts, **kwargs):
|
||||
map_tensor = th.tensor(self.timestep_map, device=ts.device, dtype=ts.dtype)
|
||||
new_ts = map_tensor[ts]
|
||||
# if self.rescale_timesteps:
|
||||
# new_ts = new_ts.float() * (1000.0 / self.original_num_steps)
|
||||
return self.model(x, new_ts, **kwargs)
|
||||
Executable
+143
@@ -0,0 +1,143 @@
|
||||
# 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 abc import ABC, abstractmethod
|
||||
|
||||
import numpy as np
|
||||
import torch as th
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
def create_named_schedule_sampler(name, diffusion):
|
||||
"""
|
||||
Create a ScheduleSampler from a library of pre-defined samplers.
|
||||
:param name: the name of the sampler.
|
||||
:param diffusion: the diffusion object to sample for.
|
||||
"""
|
||||
if name == "uniform":
|
||||
return UniformSampler(diffusion)
|
||||
elif name == "loss-second-moment":
|
||||
return LossSecondMomentResampler(diffusion)
|
||||
else:
|
||||
raise NotImplementedError(f"unknown schedule sampler: {name}")
|
||||
|
||||
|
||||
class ScheduleSampler(ABC):
|
||||
"""
|
||||
A distribution over timesteps in the diffusion process, intended to reduce
|
||||
variance of the objective.
|
||||
By default, samplers perform unbiased importance sampling, in which the
|
||||
objective's mean is unchanged.
|
||||
However, subclasses may override sample() to change how the resampled
|
||||
terms are reweighted, allowing for actual changes in the objective.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def weights(self):
|
||||
"""
|
||||
Get a numpy array of weights, one per diffusion step.
|
||||
The weights needn't be normalized, but must be positive.
|
||||
"""
|
||||
|
||||
def sample(self, batch_size, device):
|
||||
"""
|
||||
Importance-sample timesteps for a batch.
|
||||
:param batch_size: the number of timesteps.
|
||||
:param device: the torch device to save to.
|
||||
:return: a tuple (timesteps, weights):
|
||||
- timesteps: a tensor of timestep indices.
|
||||
- weights: a tensor of weights to scale the resulting losses.
|
||||
"""
|
||||
w = self.weights()
|
||||
p = w / np.sum(w)
|
||||
indices_np = np.random.choice(len(p), size=(batch_size,), p=p)
|
||||
indices = th.from_numpy(indices_np).long().to(device)
|
||||
weights_np = 1 / (len(p) * p[indices_np])
|
||||
weights = th.from_numpy(weights_np).float().to(device)
|
||||
return indices, weights
|
||||
|
||||
|
||||
class UniformSampler(ScheduleSampler):
|
||||
def __init__(self, diffusion):
|
||||
self.diffusion = diffusion
|
||||
self._weights = np.ones([diffusion.num_timesteps])
|
||||
|
||||
def weights(self):
|
||||
return self._weights
|
||||
|
||||
|
||||
class LossAwareSampler(ScheduleSampler):
|
||||
def update_with_local_losses(self, local_ts, local_losses):
|
||||
"""
|
||||
Update the reweighting using losses from a model.
|
||||
Call this method from each rank with a batch of timesteps and the
|
||||
corresponding losses for each of those timesteps.
|
||||
This method will perform synchronization to make sure all of the ranks
|
||||
maintain the exact same reweighting.
|
||||
:param local_ts: an integer Tensor of timesteps.
|
||||
:param local_losses: a 1D Tensor of losses.
|
||||
"""
|
||||
batch_sizes = [th.tensor([0], dtype=th.int32, device=local_ts.device) for _ in range(dist.get_world_size())]
|
||||
dist.all_gather(
|
||||
batch_sizes,
|
||||
th.tensor([len(local_ts)], dtype=th.int32, device=local_ts.device),
|
||||
)
|
||||
|
||||
# Pad all_gather batches to be the maximum batch size.
|
||||
batch_sizes = [x.item() for x in batch_sizes]
|
||||
max_bs = max(batch_sizes)
|
||||
|
||||
timestep_batches = [th.zeros(max_bs).to(local_ts) for bs in batch_sizes]
|
||||
loss_batches = [th.zeros(max_bs).to(local_losses) for bs in batch_sizes]
|
||||
dist.all_gather(timestep_batches, local_ts)
|
||||
dist.all_gather(loss_batches, local_losses)
|
||||
timesteps = [x.item() for y, bs in zip(timestep_batches, batch_sizes) for x in y[:bs]]
|
||||
losses = [x.item() for y, bs in zip(loss_batches, batch_sizes) for x in y[:bs]]
|
||||
self.update_with_all_losses(timesteps, losses)
|
||||
|
||||
@abstractmethod
|
||||
def update_with_all_losses(self, ts, losses):
|
||||
"""
|
||||
Update the reweighting using losses from a model.
|
||||
Sub-classes should override this method to update the reweighting
|
||||
using losses from the model.
|
||||
This method directly updates the reweighting without synchronizing
|
||||
between workers. It is called by update_with_local_losses from all
|
||||
ranks with identical arguments. Thus, it should have deterministic
|
||||
behavior to maintain state across workers.
|
||||
:param ts: a list of int timesteps.
|
||||
:param losses: a list of float losses, one per timestep.
|
||||
"""
|
||||
|
||||
|
||||
class LossSecondMomentResampler(LossAwareSampler):
|
||||
def __init__(self, diffusion, history_per_term=10, uniform_prob=0.001):
|
||||
self.diffusion = diffusion
|
||||
self.history_per_term = history_per_term
|
||||
self.uniform_prob = uniform_prob
|
||||
self._loss_history = np.zeros([diffusion.num_timesteps, history_per_term], dtype=np.float64)
|
||||
self._loss_counts = np.zeros([diffusion.num_timesteps], dtype=np.int)
|
||||
|
||||
def weights(self):
|
||||
if not self._warmed_up():
|
||||
return np.ones([self.diffusion.num_timesteps], dtype=np.float64)
|
||||
weights = np.sqrt(np.mean(self._loss_history**2, axis=-1))
|
||||
weights /= np.sum(weights)
|
||||
weights *= 1 - self.uniform_prob
|
||||
weights += self.uniform_prob / len(weights)
|
||||
return weights
|
||||
|
||||
def update_with_all_losses(self, ts, losses):
|
||||
for t, loss in zip(ts, losses):
|
||||
if self._loss_counts[t] == self.history_per_term:
|
||||
# Shift out the oldest loss term.
|
||||
self._loss_history[t, :-1] = self._loss_history[t, 1:]
|
||||
self._loss_history[t, -1] = loss
|
||||
else:
|
||||
self._loss_history[t, self._loss_counts[t]] = loss
|
||||
self._loss_counts[t] += 1
|
||||
|
||||
def _warmed_up(self):
|
||||
return (self._loss_counts == self.history_per_term).all()
|
||||
Executable
Executable
+63
@@ -0,0 +1,63 @@
|
||||
import torch
|
||||
import triton
|
||||
|
||||
from .k_fused_modulate import _modulate_bwd, _modulate_fwd
|
||||
|
||||
|
||||
class _FusedModulate(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x, scale, shift):
|
||||
y = torch.empty_like(x)
|
||||
batch, seq_len, dim = x.shape
|
||||
M = batch * seq_len
|
||||
N = dim
|
||||
x = x.view(-1, dim).contiguous()
|
||||
scale = scale.view(-1, dim).contiguous()
|
||||
shift = shift.view(-1, dim).contiguous()
|
||||
|
||||
def grid(meta):
|
||||
return (
|
||||
triton.cdiv(batch * seq_len, meta["BLOCK_M"]),
|
||||
triton.cdiv(dim, meta["BLOCK_N"]),
|
||||
)
|
||||
|
||||
_modulate_fwd[grid](x, y, scale, shift, x.stride(0), scale.stride(0), M, N, seq_len)
|
||||
|
||||
ctx.save_for_backward(x, scale)
|
||||
ctx.batch = batch
|
||||
ctx.seq_len = seq_len
|
||||
ctx.dim = dim
|
||||
return y
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, dy): # pragma: no cover # this is covered, but called directly from C++
|
||||
x, scale = ctx.saved_tensors
|
||||
|
||||
batch, seq_len, dim = ctx.batch, ctx.seq_len, ctx.dim
|
||||
M = batch * seq_len
|
||||
N = dim
|
||||
|
||||
# allocate output
|
||||
dy = dy.contiguous()
|
||||
dx = torch.empty_like(dy)
|
||||
dscale = torch.empty_like(dy)
|
||||
dshift = torch.sum(dy, dim=1)
|
||||
|
||||
def grid(meta):
|
||||
return (
|
||||
triton.cdiv(batch * seq_len, meta["BLOCK_M"]),
|
||||
triton.cdiv(dim, meta["BLOCK_N"]),
|
||||
)
|
||||
|
||||
_modulate_bwd[grid](dx, x, dy, scale, dscale, x.stride(0), scale.stride(0), M, N, seq_len)
|
||||
|
||||
dscale = torch.sum(dscale, dim=1)
|
||||
return dx, dscale, dshift
|
||||
|
||||
|
||||
def fused_modulate(
|
||||
x: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return _FusedModulate.apply(x, scale, shift)
|
||||
Executable
+102
@@ -0,0 +1,102 @@
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
CONFIG_LIST = [
|
||||
triton.Config({"BLOCK_M": 256, "BLOCK_N": 32}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 128, "BLOCK_N": 64}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 128, "BLOCK_N": 32}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 64, "BLOCK_N": 128}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 64, "BLOCK_N": 64}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 64, "BLOCK_N": 32}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 32, "BLOCK_N": 64}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 32, "BLOCK_N": 128}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 32, "BLOCK_N": 256}, num_stages=2, num_warps=4),
|
||||
]
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
configs=CONFIG_LIST,
|
||||
key=["M", "N"],
|
||||
)
|
||||
@triton.jit
|
||||
def _modulate_fwd(
|
||||
x_ptr, # *Pointer* to first input vector.
|
||||
output_ptr, # *Pointer* to output vector.
|
||||
scale_ptr,
|
||||
shift_ptr,
|
||||
m_stride,
|
||||
s_stride,
|
||||
M,
|
||||
N,
|
||||
seq_len,
|
||||
BLOCK_M: tl.constexpr, # Number of elements each program should process.
|
||||
BLOCK_N: tl.constexpr,
|
||||
# NOTE: `constexpr` so it can be used as a shape value.
|
||||
):
|
||||
row_id = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0.
|
||||
rows = row_id * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
s_rows = (row_id // seq_len) * BLOCK_M
|
||||
col_id = tl.program_id(axis=1)
|
||||
cols = col_id * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||
|
||||
x_ptrs = x_ptr + rows[:, None] * m_stride + cols[None, :]
|
||||
scale_ptrs = scale_ptr + s_rows * s_stride + cols[None, :]
|
||||
shift_ptrs = shift_ptr + s_rows * s_stride + cols[None, :]
|
||||
|
||||
col_mask = cols[None, :] < N
|
||||
block_mask = (rows[:, None] < M) & col_mask
|
||||
s_block_mask = col_mask
|
||||
x = tl.load(x_ptrs, mask=block_mask, other=0.0)
|
||||
scale = tl.load(scale_ptrs, mask=s_block_mask, other=0.0)
|
||||
shift = tl.load(shift_ptrs, mask=s_block_mask, other=0.0)
|
||||
|
||||
output = x * (1 + scale) + shift
|
||||
# Write x + y back to DRAM.
|
||||
tl.store(output_ptr + rows[:, None] * m_stride + cols[None, :], output, mask=block_mask)
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
configs=CONFIG_LIST,
|
||||
key=["M", "N"],
|
||||
)
|
||||
@triton.jit
|
||||
def _modulate_bwd(
|
||||
dx_ptr, # *Pointer* to first input vector.
|
||||
x_ptr,
|
||||
dy_ptr, # *Pointer* to output vector.
|
||||
scale_ptr,
|
||||
dscale_ptr,
|
||||
m_stride,
|
||||
s_stride,
|
||||
M,
|
||||
N,
|
||||
seq_len,
|
||||
BLOCK_M: tl.constexpr, # Number of elements each program should process.
|
||||
BLOCK_N: tl.constexpr,
|
||||
# NOTE: `constexpr` so it can be used as a shape value.
|
||||
):
|
||||
row_id = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0.
|
||||
rows = row_id * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
s_rows = (row_id // seq_len) * BLOCK_M
|
||||
col_id = tl.program_id(axis=1)
|
||||
cols = col_id * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||
|
||||
x_ptrs = x_ptr + rows[:, None] * m_stride + cols[None, :]
|
||||
dy_ptrs = dy_ptr + rows[:, None] * m_stride + cols[None, :]
|
||||
dx_ptrs = dx_ptr + rows[:, None] * m_stride + cols[None, :]
|
||||
dscale_ptrs = dscale_ptr + rows[:, None] * m_stride + cols[None, :]
|
||||
|
||||
scale_ptrs = scale_ptr + s_rows * s_stride + cols[None, :]
|
||||
|
||||
col_mask = cols[None, :] < N
|
||||
block_mask = (rows[:, None] < M) & col_mask
|
||||
s_block_mask = col_mask
|
||||
x = tl.load(x_ptrs, mask=block_mask, other=0.0)
|
||||
dy = tl.load(dy_ptrs, mask=block_mask, other=0.0)
|
||||
scale = tl.load(scale_ptrs, mask=s_block_mask, other=0.0)
|
||||
|
||||
dx = dy * (1 + scale)
|
||||
dscale = dy * x
|
||||
# Write x + y back to DRAM.
|
||||
tl.store(dx_ptrs, dx, mask=block_mask)
|
||||
tl.store(dscale_ptrs, dscale, mask=block_mask)
|
||||
Executable
Executable
+6
@@ -0,0 +1,6 @@
|
||||
from .dit import DiT, DiT_models
|
||||
|
||||
__all__ = [
|
||||
"DiT",
|
||||
"DiT_models",
|
||||
]
|
||||
Executable
+305
@@ -0,0 +1,305 @@
|
||||
# Modified from Meta DiT
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# DiT: https://github.com/facebookresearch/DiT/tree/main
|
||||
# GLIDE: https://github.com/openai/glide-text2im
|
||||
# MAE: https://github.com/facebookresearch/mae/blob/main/models_mae.py
|
||||
# --------------------------------------------------------
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.utils.checkpoint
|
||||
from timm.models.vision_transformer import Mlp, PatchEmbed
|
||||
|
||||
from opendit.modules.attn import Attention
|
||||
from opendit.modules.embed import LabelEmbedder, TimestepEmbedder, get_2d_sincos_pos_embed
|
||||
from opendit.modules.layers import FinalLayer, get_layernorm, modulate
|
||||
|
||||
|
||||
class DiTBlock(nn.Module):
|
||||
"""
|
||||
A DiT block with adaptive layer norm zero (adaLN-Zero) conditioning.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
num_heads,
|
||||
mlp_ratio=4.0,
|
||||
enable_flashattn=False,
|
||||
enable_layernorm_kernel=False,
|
||||
enable_modulate_kernel=False,
|
||||
**block_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.enable_modulate_kernel = enable_modulate_kernel
|
||||
self.norm1 = get_layernorm(hidden_size, eps=1e-6, affine=False, use_kernel=enable_layernorm_kernel)
|
||||
self.attn = Attention(
|
||||
hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
enable_flashattn=enable_flashattn,
|
||||
**block_kwargs,
|
||||
)
|
||||
self.norm2 = get_layernorm(hidden_size, eps=1e-6, affine=False, use_kernel=enable_layernorm_kernel)
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
approx_gelu = lambda: nn.GELU(approximate="tanh")
|
||||
self.mlp = Mlp(in_features=hidden_size, hidden_features=mlp_hidden_dim, act_layer=approx_gelu, drop=0)
|
||||
self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True))
|
||||
|
||||
def forward(self, x, c):
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1)
|
||||
x = x + gate_msa.unsqueeze(1) * self.attn(
|
||||
modulate(self.norm1, x, shift_msa, scale_msa, self.enable_modulate_kernel)
|
||||
)
|
||||
x = x + gate_mlp.unsqueeze(1) * self.mlp(
|
||||
modulate(self.norm2, x, shift_mlp, scale_mlp, self.enable_modulate_kernel)
|
||||
)
|
||||
return x
|
||||
|
||||
|
||||
class DiT(nn.Module):
|
||||
"""
|
||||
Diffusion model with a Transformer backbone.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_size=32,
|
||||
patch_size=2,
|
||||
in_channels=4,
|
||||
hidden_size=1152,
|
||||
depth=28,
|
||||
num_heads=16,
|
||||
mlp_ratio=4.0,
|
||||
class_dropout_prob=0.1,
|
||||
num_classes=1000,
|
||||
learn_sigma: bool = True,
|
||||
enable_flashattn: bool = False,
|
||||
enable_layernorm_kernel: bool = False,
|
||||
enable_modulate_kernel: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
):
|
||||
super().__init__()
|
||||
self.learn_sigma = learn_sigma
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = in_channels * 2 if learn_sigma else in_channels
|
||||
self.hidden_size = hidden_size
|
||||
self.patch_size = patch_size
|
||||
self.input_size = input_size
|
||||
self.num_heads = num_heads
|
||||
self.dtype = dtype
|
||||
if enable_flashattn:
|
||||
assert dtype in [
|
||||
torch.float16,
|
||||
torch.bfloat16,
|
||||
], f"Flash attention only supports float16 and bfloat16, but got {self.dtype}"
|
||||
|
||||
self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size, bias=True)
|
||||
self.t_embedder = TimestepEmbedder(hidden_size)
|
||||
self.y_embedder = LabelEmbedder(num_classes, hidden_size, class_dropout_prob)
|
||||
self.num_patches = self.x_embedder.num_patches
|
||||
|
||||
self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches, hidden_size), requires_grad=False)
|
||||
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
DiTBlock(
|
||||
hidden_size,
|
||||
num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
enable_flashattn=enable_flashattn,
|
||||
enable_modulate_kernel=enable_modulate_kernel,
|
||||
enable_layernorm_kernel=enable_layernorm_kernel,
|
||||
)
|
||||
for _ in range(depth)
|
||||
]
|
||||
)
|
||||
self.final_layer = FinalLayer(hidden_size, patch_size**2, self.out_channels)
|
||||
self.initialize_weights()
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
self.gradient_checkpointing = True
|
||||
|
||||
def initialize_weights(self):
|
||||
# Initialize transformer layers:
|
||||
def _basic_init(module):
|
||||
if isinstance(module, nn.Linear):
|
||||
if module.weight.requires_grad:
|
||||
torch.nn.init.xavier_uniform_(module.weight)
|
||||
if module.bias is not None:
|
||||
nn.init.constant_(module.bias, 0)
|
||||
|
||||
self.apply(_basic_init)
|
||||
|
||||
# Initialize (and freeze) pos_embed by sin-cos embedding:
|
||||
pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], int(self.x_embedder.num_patches**0.5))
|
||||
self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0))
|
||||
|
||||
# Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
|
||||
w = self.x_embedder.proj.weight.data
|
||||
nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
|
||||
nn.init.constant_(self.x_embedder.proj.bias, 0)
|
||||
|
||||
# Initialize label embedding table:
|
||||
if isinstance(self.y_embedder, LabelEmbedder):
|
||||
nn.init.normal_(self.y_embedder.embedding_table.weight, std=0.02)
|
||||
|
||||
# Initialize timestep embedding MLP:
|
||||
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
|
||||
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
|
||||
|
||||
# Zero-out adaLN modulation layers in DiT blocks:
|
||||
for block in self.blocks:
|
||||
nn.init.constant_(block.adaLN_modulation[-1].weight, 0)
|
||||
nn.init.constant_(block.adaLN_modulation[-1].bias, 0)
|
||||
|
||||
# Zero-out output layers:
|
||||
nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)
|
||||
nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)
|
||||
nn.init.constant_(self.final_layer.linear.weight, 0)
|
||||
nn.init.constant_(self.final_layer.linear.bias, 0)
|
||||
|
||||
def unpatchify(self, x):
|
||||
"""
|
||||
x: (N, T, patch_size**2 * C)
|
||||
imgs: (N, H, W, C)
|
||||
"""
|
||||
c = self.out_channels
|
||||
p = self.x_embedder.patch_size[0]
|
||||
h = w = int(x.shape[1] ** 0.5)
|
||||
assert h * w == x.shape[1]
|
||||
|
||||
x = x.reshape(shape=(x.shape[0], h, w, p, p, c))
|
||||
x = torch.einsum("nhwpqc->nchpwq", x)
|
||||
imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p))
|
||||
return imgs
|
||||
|
||||
@staticmethod
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
def forward(self, x, t, y):
|
||||
"""
|
||||
Forward pass of DiT.
|
||||
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||
t: (N,) tensor of diffusion timesteps
|
||||
y: (N,) tensor of class labels
|
||||
"""
|
||||
|
||||
# origin inputs should be float32, cast to specified dtype
|
||||
x = x.to(self.dtype)
|
||||
|
||||
x = self.x_embedder(x) + self.pos_embed # (N, T, D), where T = H * W / patch_size ** 2
|
||||
|
||||
t = self.t_embedder(t, dtype=x.dtype) # (N, D)
|
||||
y = self.y_embedder(y, self.training) # (N, D)
|
||||
c = t + y # (N, D)
|
||||
|
||||
for block in self.blocks:
|
||||
if self.gradient_checkpointing:
|
||||
x = torch.utils.checkpoint.checkpoint(self.create_custom_forward(block), x, c)
|
||||
else:
|
||||
x = block(x, c) # (N, T, D)
|
||||
|
||||
x = self.final_layer(x, c) # (N, T, patch_size ** 2 * out_channels)
|
||||
x = self.unpatchify(x) # (N, out_channels, H, W)
|
||||
|
||||
# cast to float32 for better accuracy
|
||||
x = x.to(torch.float32)
|
||||
return x
|
||||
|
||||
def forward_with_cfg(self, x, t, y, cfg_scale):
|
||||
"""
|
||||
Forward pass of DiT, but also batches the unconditional forward pass for classifier-free guidance.
|
||||
"""
|
||||
# https://github.com/openai/glide-text2im/blob/main/notebooks/text2im.ipynb
|
||||
half = x[: len(x) // 2]
|
||||
combined = torch.cat([half, half], dim=0)
|
||||
model_out = self.forward(combined, t, y)
|
||||
# For exact reproducibility reasons, we apply classifier-free guidance on only
|
||||
# three channels by default. The standard approach to cfg applies it to all channels.
|
||||
# This can be done by uncommenting the following line and commenting-out the line following that.
|
||||
# eps, rest = model_out[:, :self.in_channels], model_out[:, self.in_channels:]
|
||||
eps, rest = model_out[:, :3], model_out[:, 3:]
|
||||
cond_eps, uncond_eps = torch.split(eps, len(eps) // 2, dim=0)
|
||||
half_eps = uncond_eps + cfg_scale * (cond_eps - uncond_eps)
|
||||
eps = torch.cat([half_eps, half_eps], dim=0)
|
||||
return torch.cat([eps, rest], dim=1)
|
||||
|
||||
|
||||
#################################################################################
|
||||
# DiT Configs #
|
||||
#################################################################################
|
||||
|
||||
|
||||
def DiT_XL_2(**kwargs):
|
||||
return DiT(depth=28, hidden_size=1152, patch_size=2, num_heads=16, **kwargs)
|
||||
|
||||
|
||||
def DiT_XL_4(**kwargs):
|
||||
return DiT(depth=28, hidden_size=1152, patch_size=4, num_heads=16, **kwargs)
|
||||
|
||||
|
||||
def DiT_XL_8(**kwargs):
|
||||
return DiT(depth=28, hidden_size=1152, patch_size=8, num_heads=16, **kwargs)
|
||||
|
||||
|
||||
def DiT_L_2(**kwargs):
|
||||
return DiT(depth=24, hidden_size=1024, patch_size=2, num_heads=16, **kwargs)
|
||||
|
||||
|
||||
def DiT_L_4(**kwargs):
|
||||
return DiT(depth=24, hidden_size=1024, patch_size=4, num_heads=16, **kwargs)
|
||||
|
||||
|
||||
def DiT_L_8(**kwargs):
|
||||
return DiT(depth=24, hidden_size=1024, patch_size=8, num_heads=16, **kwargs)
|
||||
|
||||
|
||||
def DiT_B_2(**kwargs):
|
||||
return DiT(depth=12, hidden_size=768, patch_size=2, num_heads=12, **kwargs)
|
||||
|
||||
|
||||
def DiT_B_4(**kwargs):
|
||||
return DiT(depth=12, hidden_size=768, patch_size=4, num_heads=12, **kwargs)
|
||||
|
||||
|
||||
def DiT_B_8(**kwargs):
|
||||
return DiT(depth=12, hidden_size=768, patch_size=8, num_heads=12, **kwargs)
|
||||
|
||||
|
||||
def DiT_S_2(**kwargs):
|
||||
return DiT(depth=12, hidden_size=384, patch_size=2, num_heads=6, **kwargs)
|
||||
|
||||
|
||||
def DiT_S_4(**kwargs):
|
||||
return DiT(depth=12, hidden_size=384, patch_size=4, num_heads=6, **kwargs)
|
||||
|
||||
|
||||
def DiT_S_8(**kwargs):
|
||||
return DiT(depth=12, hidden_size=384, patch_size=8, num_heads=6, **kwargs)
|
||||
|
||||
|
||||
DiT_models = {
|
||||
"DiT-XL/2": DiT_XL_2,
|
||||
"DiT-XL/4": DiT_XL_4,
|
||||
"DiT-XL/8": DiT_XL_8,
|
||||
"DiT-L/2": DiT_L_2,
|
||||
"DiT-L/4": DiT_L_4,
|
||||
"DiT-L/8": DiT_L_8,
|
||||
"DiT-B/2": DiT_B_2,
|
||||
"DiT-B/4": DiT_B_4,
|
||||
"DiT-B/8": DiT_B_8,
|
||||
"DiT-S/2": DiT_S_2,
|
||||
"DiT-S/4": DiT_S_4,
|
||||
"DiT-S/8": DiT_S_8,
|
||||
}
|
||||
Executable
+7
@@ -0,0 +1,7 @@
|
||||
from .latte_t2v import LatteT2V
|
||||
from .pipeline import LattePipeline
|
||||
|
||||
__all__ = [
|
||||
"LatteT2V",
|
||||
"LattePipeline",
|
||||
]
|
||||
Executable
+1399
File diff suppressed because it is too large
Load Diff
Executable
+827
@@ -0,0 +1,827 @@
|
||||
# Adapted from Latte
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Latte: https://github.com/Vchitect/Latte
|
||||
# --------------------------------------------------------
|
||||
|
||||
import html
|
||||
import inspect
|
||||
import re
|
||||
import urllib.parse as ul
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, List, Optional, Tuple, Union
|
||||
|
||||
import einops
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
from diffusers.models import AutoencoderKL, Transformer2DModel
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.schedulers import DPMSolverMultistepScheduler
|
||||
from diffusers.utils import (
|
||||
BACKENDS_MAPPING,
|
||||
BaseOutput,
|
||||
is_bs4_available,
|
||||
is_ftfy_available,
|
||||
logging,
|
||||
replace_example_docstring,
|
||||
)
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from transformers import T5EncoderModel, T5Tokenizer
|
||||
|
||||
from opendit.core.pab_mgr import get_diffusion_skip, get_diffusion_skip_timestep, skip_diffusion_timestep
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
if is_bs4_available():
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
if is_ftfy_available():
|
||||
import ftfy
|
||||
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
Examples:
|
||||
```py
|
||||
>>> import torch
|
||||
>>> from diffusers import PixArtAlphaPipeline
|
||||
|
||||
>>> # You can replace the checkpoint id with "PixArt-alpha/PixArt-XL-2-512x512" too.
|
||||
>>> pipe = PixArtAlphaPipeline.from_pretrained("PixArt-alpha/PixArt-XL-2-1024-MS", torch_dtype=torch.float16)
|
||||
>>> # Enable memory optimizations.
|
||||
>>> pipe.enable_model_cpu_offload()
|
||||
|
||||
>>> prompt = "A small cactus with a happy face in the Sahara desert."
|
||||
>>> image = pipe(prompt).images[0]
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoPipelineOutput(BaseOutput):
|
||||
video: torch.Tensor
|
||||
|
||||
|
||||
class LattePipeline(DiffusionPipeline):
|
||||
r"""
|
||||
Pipeline for text-to-image generation using PixArt-Alpha.
|
||||
|
||||
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.)
|
||||
|
||||
Args:
|
||||
vae ([`AutoencoderKL`]):
|
||||
Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.
|
||||
text_encoder ([`T5EncoderModel`]):
|
||||
Frozen text-encoder. PixArt-Alpha uses
|
||||
[T5](https://huggingface.co/docs/transformers/model_doc/t5#transformers.T5EncoderModel), specifically the
|
||||
[t5-v1_1-xxl](https://huggingface.co/PixArt-alpha/PixArt-alpha/tree/main/t5-v1_1-xxl) variant.
|
||||
tokenizer (`T5Tokenizer`):
|
||||
Tokenizer of class
|
||||
[T5Tokenizer](https://huggingface.co/docs/transformers/model_doc/t5#transformers.T5Tokenizer).
|
||||
transformer ([`Transformer2DModel`]):
|
||||
A text conditioned `Transformer2DModel` to denoise the encoded image latents.
|
||||
scheduler ([`SchedulerMixin`]):
|
||||
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
|
||||
"""
|
||||
bad_punct_regex = re.compile(
|
||||
r"[" + "#®•©™&@·º½¾¿¡§~" + "\)" + "\(" + "\]" + "\[" + "\}" + "\{" + "\|" + "\\" + "\/" + "\*" + r"]{1,}"
|
||||
) # noqa
|
||||
|
||||
_optional_components = ["tokenizer", "text_encoder"]
|
||||
model_cpu_offload_seq = "text_encoder->transformer->vae"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: T5Tokenizer,
|
||||
text_encoder: T5EncoderModel,
|
||||
vae: AutoencoderKL,
|
||||
transformer: Transformer2DModel,
|
||||
scheduler: DPMSolverMultistepScheduler,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
tokenizer=tokenizer, text_encoder=text_encoder, vae=vae, transformer=transformer, scheduler=scheduler
|
||||
)
|
||||
|
||||
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:
|
||||
keep_index = mask.sum().item()
|
||||
return emb[:, :, :keep_index, :], keep_index # 1, 120, 4096 -> 1 7 4096
|
||||
else:
|
||||
masked_feature = emb * mask[:, None, :, None] # 1 120 4096
|
||||
return masked_feature, emb.shape[2]
|
||||
|
||||
# Adapted from diffusers.pipelines.deepfloyd_if.pipeline_if.encode_prompt
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
do_classifier_free_guidance: bool = True,
|
||||
negative_prompt: str = "",
|
||||
num_images_per_prompt: int = 1,
|
||||
device: Optional[torch.device] = None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
clean_caption: bool = False,
|
||||
mask_feature: bool = True,
|
||||
):
|
||||
r"""
|
||||
Encodes the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
prompt to be encoded
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt 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`). For
|
||||
PixArt-Alpha, this should be "".
|
||||
do_classifier_free_guidance (`bool`, *optional*, defaults to `True`):
|
||||
whether to use classifier free guidance or not
|
||||
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||
number of images that should be generated per prompt
|
||||
device: (`torch.device`, *optional*):
|
||||
torch device to place the resulting embeddings on
|
||||
prompt_embeds (`torch.FloatTensor`, *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.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`):
|
||||
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.
|
||||
"""
|
||||
embeds_initially_provided = prompt_embeds is not None and negative_prompt_embeds is not None
|
||||
|
||||
if device is None:
|
||||
device = self._execution_device
|
||||
|
||||
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]
|
||||
|
||||
# See Section 3.1. of the paper.
|
||||
max_length = 120
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt = self._text_preprocessing(prompt, clean_caption=clean_caption)
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_length,
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
untruncated_ids = self.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
|
||||
):
|
||||
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, max_length - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because CLIP can only handle sequences up to"
|
||||
f" {max_length} tokens: {removed_text}"
|
||||
)
|
||||
|
||||
attention_mask = text_inputs.attention_mask.to(device)
|
||||
prompt_embeds_attention_mask = attention_mask
|
||||
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=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
|
||||
elif self.transformer is not None:
|
||||
dtype = self.transformer.dtype
|
||||
else:
|
||||
dtype = None
|
||||
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
bs_embed, seq_len, _ = prompt_embeds.shape
|
||||
# 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)
|
||||
|
||||
# get unconditional embeddings for classifier free guidance
|
||||
if do_classifier_free_guidance and negative_prompt_embeds is None:
|
||||
uncond_tokens = [negative_prompt] * batch_size
|
||||
uncond_tokens = self._text_preprocessing(uncond_tokens, clean_caption=clean_caption)
|
||||
max_length = prompt_embeds.shape[1]
|
||||
uncond_input = self.tokenizer(
|
||||
uncond_tokens,
|
||||
padding="max_length",
|
||||
max_length=max_length,
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
attention_mask = uncond_input.attention_mask.to(device)
|
||||
|
||||
negative_prompt_embeds = self.text_encoder(
|
||||
uncond_input.input_ids.to(device),
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
negative_prompt_embeds = negative_prompt_embeds[0]
|
||||
|
||||
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)
|
||||
|
||||
# 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
|
||||
else:
|
||||
negative_prompt_embeds = None
|
||||
|
||||
# Perform additional masking.
|
||||
if mask_feature and not embeds_initially_provided:
|
||||
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
|
||||
)
|
||||
|
||||
# import torch.nn.functional as F
|
||||
|
||||
# padding = (0, 0, 0, 113) # (左, 右, 下, 上)
|
||||
# masked_prompt_embeds_ = F.pad(masked_prompt_embeds, padding, "constant", 0)
|
||||
# masked_negative_prompt_embeds_ = F.pad(masked_negative_prompt_embeds, padding, "constant", 0)
|
||||
|
||||
# print(masked_prompt_embeds == masked_prompt_embeds_[:, :masked_negative_prompt_embeds.shape[1], ...])
|
||||
|
||||
return masked_prompt_embeds, masked_negative_prompt_embeds
|
||||
# return masked_prompt_embeds_, masked_negative_prompt_embeds_
|
||||
|
||||
return prompt_embeds, negative_prompt_embeds
|
||||
|
||||
# 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,
|
||||
callback_steps,
|
||||
prompt_embeds=None,
|
||||
negative_prompt_embeds=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_steps is None) or (
|
||||
callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0)
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_steps` has to be a positive integer but is {callback_steps} of type"
|
||||
f" {type(callback_steps)}."
|
||||
)
|
||||
|
||||
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 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 is not None and negative_prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `negative_prompt_embeds`:"
|
||||
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
|
||||
)
|
||||
|
||||
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 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}."
|
||||
)
|
||||
|
||||
# Copied from diffusers.pipelines.deepfloyd_if.pipeline_if.IFPipeline._text_preprocessing
|
||||
def _text_preprocessing(self, text, clean_caption=False):
|
||||
if clean_caption and not is_bs4_available():
|
||||
logger.warn(BACKENDS_MAPPING["bs4"][-1].format("Setting `clean_caption=True`"))
|
||||
logger.warn("Setting `clean_caption` to False...")
|
||||
clean_caption = False
|
||||
|
||||
if clean_caption and not is_ftfy_available():
|
||||
logger.warn(BACKENDS_MAPPING["ftfy"][-1].format("Setting `clean_caption=True`"))
|
||||
logger.warn("Setting `clean_caption` to False...")
|
||||
clean_caption = False
|
||||
|
||||
if not isinstance(text, (tuple, list)):
|
||||
text = [text]
|
||||
|
||||
def process(text: str):
|
||||
if clean_caption:
|
||||
text = self._clean_caption(text)
|
||||
text = self._clean_caption(text)
|
||||
else:
|
||||
text = text.lower().strip()
|
||||
return text
|
||||
|
||||
return [process(t) for t in text]
|
||||
|
||||
# Copied from diffusers.pipelines.deepfloyd_if.pipeline_if.IFPipeline._clean_caption
|
||||
def _clean_caption(self, caption):
|
||||
caption = str(caption)
|
||||
caption = ul.unquote_plus(caption)
|
||||
caption = caption.strip().lower()
|
||||
caption = re.sub("<person>", "person", caption)
|
||||
# urls:
|
||||
caption = re.sub(
|
||||
r"\b((?:https?:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))", # noqa
|
||||
"",
|
||||
caption,
|
||||
) # regex for urls
|
||||
caption = re.sub(
|
||||
r"\b((?:www:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))", # noqa
|
||||
"",
|
||||
caption,
|
||||
) # regex for urls
|
||||
# html:
|
||||
caption = BeautifulSoup(caption, features="html.parser").text
|
||||
|
||||
# @<nickname>
|
||||
caption = re.sub(r"@[\w\d]+\b", "", caption)
|
||||
|
||||
# 31C0—31EF CJK Strokes
|
||||
# 31F0—31FF Katakana Phonetic Extensions
|
||||
# 3200—32FF Enclosed CJK Letters and Months
|
||||
# 3300—33FF CJK Compatibility
|
||||
# 3400—4DBF CJK Unified Ideographs Extension A
|
||||
# 4DC0—4DFF Yijing Hexagram Symbols
|
||||
# 4E00—9FFF CJK Unified Ideographs
|
||||
caption = re.sub(r"[\u31c0-\u31ef]+", "", caption)
|
||||
caption = re.sub(r"[\u31f0-\u31ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3200-\u32ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3300-\u33ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3400-\u4dbf]+", "", caption)
|
||||
caption = re.sub(r"[\u4dc0-\u4dff]+", "", caption)
|
||||
caption = re.sub(r"[\u4e00-\u9fff]+", "", caption)
|
||||
#######################################################
|
||||
|
||||
# все виды тире / all types of dash --> "-"
|
||||
caption = re.sub(
|
||||
r"[\u002D\u058A\u05BE\u1400\u1806\u2010-\u2015\u2E17\u2E1A\u2E3A\u2E3B\u2E40\u301C\u3030\u30A0\uFE31\uFE32\uFE58\uFE63\uFF0D]+", # noqa
|
||||
"-",
|
||||
caption,
|
||||
)
|
||||
|
||||
# кавычки к одному стандарту
|
||||
caption = re.sub(r"[`´«»“”¨]", '"', caption)
|
||||
caption = re.sub(r"[‘’]", "'", caption)
|
||||
|
||||
# "
|
||||
caption = re.sub(r""?", "", caption)
|
||||
# &
|
||||
caption = re.sub(r"&", "", caption)
|
||||
|
||||
# ip adresses:
|
||||
caption = re.sub(r"\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}", " ", caption)
|
||||
|
||||
# article ids:
|
||||
caption = re.sub(r"\d:\d\d\s+$", "", caption)
|
||||
|
||||
# \n
|
||||
caption = re.sub(r"\\n", " ", caption)
|
||||
|
||||
# "#123"
|
||||
caption = re.sub(r"#\d{1,3}\b", "", caption)
|
||||
# "#12345.."
|
||||
caption = re.sub(r"#\d{5,}\b", "", caption)
|
||||
# "123456.."
|
||||
caption = re.sub(r"\b\d{6,}\b", "", caption)
|
||||
# filenames:
|
||||
caption = re.sub(r"[\S]+\.(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)", "", caption)
|
||||
|
||||
#
|
||||
caption = re.sub(r"[\"\']{2,}", r'"', caption) # """AUSVERKAUFT"""
|
||||
caption = re.sub(r"[\.]{2,}", r" ", caption) # """AUSVERKAUFT"""
|
||||
|
||||
caption = re.sub(self.bad_punct_regex, r" ", caption) # ***AUSVERKAUFT***, #AUSVERKAUFT
|
||||
caption = re.sub(r"\s+\.\s+", r" ", caption) # " . "
|
||||
|
||||
# this-is-my-cute-cat / this_is_my_cute_cat
|
||||
regex2 = re.compile(r"(?:\-|\_)")
|
||||
if len(re.findall(regex2, caption)) > 3:
|
||||
caption = re.sub(regex2, " ", caption)
|
||||
|
||||
caption = ftfy.fix_text(caption)
|
||||
caption = html.unescape(html.unescape(caption))
|
||||
|
||||
caption = re.sub(r"\b[a-zA-Z]{1,3}\d{3,15}\b", "", caption) # jc6640
|
||||
caption = re.sub(r"\b[a-zA-Z]+\d+[a-zA-Z]+\b", "", caption) # jc6640vc
|
||||
caption = re.sub(r"\b\d+[a-zA-Z]+\d+\b", "", caption) # 6640vc231
|
||||
|
||||
caption = re.sub(r"(worldwide\s+)?(free\s+)?shipping", "", caption)
|
||||
caption = re.sub(r"(free\s)?download(\sfree)?", "", caption)
|
||||
caption = re.sub(r"\bclick\b\s(?:for|on)\s\w+", "", caption)
|
||||
caption = re.sub(r"\b(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)(\simage[s]?)?", "", caption)
|
||||
caption = re.sub(r"\bpage\s+\d+\b", "", caption)
|
||||
|
||||
caption = re.sub(r"\b\d*[a-zA-Z]+\d+[a-zA-Z]+\d+[a-zA-Z\d]*\b", r" ", caption) # j2d1a2a...
|
||||
|
||||
caption = re.sub(r"\b\d+\.?\d*[xх×]\d+\.?\d*\b", "", caption)
|
||||
|
||||
caption = re.sub(r"\b\s+\:\s+", r": ", caption)
|
||||
caption = re.sub(r"(\D[,\./])\b", r"\1 ", caption)
|
||||
caption = re.sub(r"\s+", " ", caption)
|
||||
|
||||
caption.strip()
|
||||
|
||||
caption = re.sub(r"^[\"\']([\w\W]+)[\"\']$", r"\1", caption)
|
||||
caption = re.sub(r"^[\'\_,\-\:;]", r"", caption)
|
||||
caption = re.sub(r"[\'\_,\-\:\-\+]$", r"", caption)
|
||||
caption = re.sub(r"^\.\S+$", "", caption)
|
||||
|
||||
return caption.strip()
|
||||
|
||||
# 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
|
||||
):
|
||||
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
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
negative_prompt: str = "",
|
||||
num_inference_steps: int = 20,
|
||||
timesteps: List[int] = None,
|
||||
guidance_scale: float = 4.5,
|
||||
num_images_per_prompt: Optional[int] = 1,
|
||||
video_length: Optional[int] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
eta: float = 0.0,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
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,
|
||||
enable_temporal_attentions: bool = True,
|
||||
enable_vae_temporal_decoder: bool = False,
|
||||
verbose: bool = False,
|
||||
) -> Union[VideoPipelineOutput, Tuple]:
|
||||
"""
|
||||
Function invoked when calling the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
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`).
|
||||
num_inference_steps (`int`, *optional*, defaults to 100):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
timesteps (`List[int]`, *optional*):
|
||||
Custom timesteps to use for the denoising process. If not defined, equal spaced `num_inference_steps`
|
||||
timesteps are used. Must be in descending order.
|
||||
guidance_scale (`float`, *optional*, defaults to 7.0):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
|
||||
`guidance_scale` is defined as `w` of equation 2. of [Imagen
|
||||
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
|
||||
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
|
||||
usually at the expense of lower image quality.
|
||||
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
height (`int`, *optional*, defaults to self.unet.config.sample_size):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, *optional*, defaults to self.unet.config.sample_size):
|
||||
The width in pixels of the generated image.
|
||||
eta (`float`, *optional*, defaults to 0.0):
|
||||
Corresponds to parameter eta (η) in the DDIM paper: https://arxiv.org/abs/2010.02502. Only applies to
|
||||
[`schedulers.DDIMScheduler`], will be ignored for others.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
|
||||
to make generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor will ge generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.FloatTensor`, *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.FloatTensor`, *optional*):
|
||||
Pre-generated negative text embeddings. For PixArt-Alpha this negative prompt should be "". If not
|
||||
provided, negative_prompt_embeds will be generated from `negative_prompt` input argument.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generate image. Choose between
|
||||
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~pipelines.stable_diffusion.IFPipelineOutput`] instead of a plain tuple.
|
||||
callback (`Callable`, *optional*):
|
||||
A function that will be called every `callback_steps` steps during inference. The function will be
|
||||
called with the following arguments: `callback(step: int, timestep: int, latents: torch.FloatTensor)`.
|
||||
callback_steps (`int`, *optional*, defaults to 1):
|
||||
The frequency at which the `callback` function will be called. If not specified, the callback will be
|
||||
called at every step.
|
||||
clean_caption (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to clean the caption before creating embeddings. Requires `beautifulsoup4` and `ftfy` to
|
||||
be installed. If the dependencies are not installed, the embeddings will be created from the raw
|
||||
prompt.
|
||||
mask_feature (`bool` defaults to `True`): If set to `True`, the text embeddings will be masked.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~pipelines.ImagePipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`, [`~pipelines.ImagePipelineOutput`] is returned, otherwise a `tuple` is
|
||||
returned where the first element is a list with the generated images
|
||||
"""
|
||||
# 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
|
||||
self.check_inputs(prompt, height, width, negative_prompt, callback_steps, prompt_embeds, negative_prompt_embeds)
|
||||
|
||||
# 2. Default height and width to transformer
|
||||
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
|
||||
|
||||
# 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.
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
|
||||
# 3. Encode input prompt
|
||||
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
|
||||
prompt,
|
||||
do_classifier_free_guidance,
|
||||
negative_prompt=negative_prompt,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
device=device,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
clean_caption=clean_caption,
|
||||
mask_feature=mask_feature,
|
||||
)
|
||||
if do_classifier_free_guidance:
|
||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
# timesteps = self.scheduler.timesteps # NOTE change timestep_respacing here
|
||||
|
||||
if get_diffusion_skip() and get_diffusion_skip_timestep() is not None:
|
||||
# TODO add assertion for timestep_respacing
|
||||
# timestep_respacing = get_diffusion_skip_timestep()
|
||||
# timesteps = space_timesteps(1000, timestep_respacing)
|
||||
|
||||
diffusion_skip_timestep = get_diffusion_skip_timestep()
|
||||
timesteps = skip_diffusion_timestep(self.scheduler.timesteps, diffusion_skip_timestep)
|
||||
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
orignal_timesteps = self.scheduler.timesteps
|
||||
|
||||
if verbose and dist.get_rank() == 0:
|
||||
print("============================")
|
||||
print("skip diffusion steps!!!")
|
||||
print("============================")
|
||||
print(f"orignal sample timesteps: {orignal_timesteps}")
|
||||
print(f"orignal diffusion steps: {len(orignal_timesteps)}")
|
||||
print("============================")
|
||||
print(f"skip diffusion steps: {get_diffusion_skip_timestep()}")
|
||||
print(f"sample timesteps: {timesteps}")
|
||||
print(f"num_inference_steps: {len(timesteps)}")
|
||||
print("============================")
|
||||
else:
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 5. Prepare latents.
|
||||
latent_channels = self.transformer.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_images_per_prompt,
|
||||
latent_channels,
|
||||
video_length,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
# 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)
|
||||
|
||||
# 6.1 Prepare micro-conditions.
|
||||
added_cond_kwargs = {"resolution": None, "aspect_ratio": None}
|
||||
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)
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
current_timestep = t
|
||||
if not torch.is_tensor(current_timestep):
|
||||
# TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can
|
||||
# This would be a good case for the `match` statement (Python 3.10+)
|
||||
is_mps = latent_model_input.device.type == "mps"
|
||||
if isinstance(current_timestep, float):
|
||||
dtype = torch.float32 if is_mps else torch.float64
|
||||
else:
|
||||
dtype = torch.int32 if is_mps else torch.int64
|
||||
current_timestep = torch.tensor([current_timestep], dtype=dtype, device=latent_model_input.device)
|
||||
elif len(current_timestep.shape) == 0:
|
||||
current_timestep = current_timestep[None].to(latent_model_input.device)
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
current_timestep = current_timestep.expand(latent_model_input.shape[0])
|
||||
|
||||
# predict noise model_output
|
||||
noise_pred = self.transformer(
|
||||
latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep=current_timestep,
|
||||
added_cond_kwargs=added_cond_kwargs,
|
||||
enable_temporal_attentions=enable_temporal_attentions,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# perform guidance
|
||||
if 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)
|
||||
|
||||
# learned sigma
|
||||
if self.transformer.config.out_channels // 2 == latent_channels:
|
||||
noise_pred = noise_pred.chunk(2, dim=1)[0]
|
||||
else:
|
||||
noise_pred = noise_pred
|
||||
|
||||
# compute previous image: x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
|
||||
# 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()
|
||||
if callback is not None and i % callback_steps == 0:
|
||||
step_idx = i // getattr(self.scheduler, "order", 1)
|
||||
callback(step_idx, t, latents)
|
||||
|
||||
if not output_type == "latents":
|
||||
if latents.shape[2] == 1: # image
|
||||
video = self.decode_latents_image(latents)
|
||||
else: # video
|
||||
if enable_vae_temporal_decoder:
|
||||
video = self.decode_latents_with_temporal_decoder(latents)
|
||||
else:
|
||||
video = self.decode_latents(latents)
|
||||
else:
|
||||
video = latents
|
||||
return VideoPipelineOutput(video=video)
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (video,)
|
||||
|
||||
return VideoPipelineOutput(video=video)
|
||||
|
||||
def decode_latents_image(self, latents):
|
||||
video_length = latents.shape[2]
|
||||
latents = 1 / self.vae.config.scaling_factor * latents
|
||||
latents = einops.rearrange(latents, "b c f h w -> (b f) c h w")
|
||||
video = []
|
||||
for frame_idx in range(latents.shape[0]):
|
||||
video.append(self.vae.decode(latents[frame_idx : frame_idx + 1]).sample)
|
||||
video = torch.cat(video)
|
||||
video = einops.rearrange(video, "(b f) c h w -> b f c h w", f=video_length)
|
||||
video = (video / 2.0 + 0.5).clamp(0, 1)
|
||||
return video
|
||||
|
||||
def decode_latents(self, latents):
|
||||
video_length = latents.shape[2]
|
||||
latents = 1 / self.vae.config.scaling_factor * latents
|
||||
latents = einops.rearrange(latents, "b c f h w -> (b f) c h w")
|
||||
video = []
|
||||
for frame_idx in range(latents.shape[0]):
|
||||
video.append(self.vae.decode(latents[frame_idx : frame_idx + 1]).sample)
|
||||
video = torch.cat(video)
|
||||
video = einops.rearrange(video, "(b f) c h w -> b f h w c", f=video_length)
|
||||
video = ((video / 2.0 + 0.5).clamp(0, 1) * 255).to(dtype=torch.uint8).cpu().contiguous()
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16
|
||||
return video
|
||||
|
||||
def decode_latents_with_temporal_decoder(self, latents):
|
||||
video_length = latents.shape[2]
|
||||
latents = 1 / self.vae.config.scaling_factor * latents
|
||||
latents = einops.rearrange(latents, "b c f h w -> (b f) c h w")
|
||||
video = []
|
||||
|
||||
decode_chunk_size = 14
|
||||
for frame_idx in range(0, latents.shape[0], decode_chunk_size):
|
||||
num_frames_in = latents[frame_idx : frame_idx + decode_chunk_size].shape[0]
|
||||
|
||||
decode_kwargs = {}
|
||||
decode_kwargs["num_frames"] = num_frames_in
|
||||
|
||||
video.append(self.vae.decode(latents[frame_idx : frame_idx + decode_chunk_size], **decode_kwargs).sample)
|
||||
|
||||
video = torch.cat(video)
|
||||
video = einops.rearrange(video, "(b f) c h w -> b f h w c", f=video_length)
|
||||
video = ((video / 2.0 + 0.5).clamp(0, 1) * 255).to(dtype=torch.uint8).cpu().contiguous()
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16
|
||||
return video
|
||||
Executable
+6
@@ -0,0 +1,6 @@
|
||||
from .rflow import RFLOW
|
||||
from .stdit3 import STDiT3_XL_2
|
||||
from .text_encoder import T5Encoder, text_preprocessing
|
||||
from .vae import OpenSoraVAE_V1_2
|
||||
|
||||
__all__ = ["RFLOW", "STDiT3_XL_2", "T5Encoder", "text_preprocessing", "OpenSoraVAE_V1_2"]
|
||||
Executable
+788
@@ -0,0 +1,788 @@
|
||||
# Adapted from OpenSora
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# OpenSora: https://github.com/hpcaitech/Open-Sora
|
||||
# --------------------------------------------------------
|
||||
|
||||
|
||||
import numbers
|
||||
import os
|
||||
import re
|
||||
|
||||
import numpy as np
|
||||
import requests
|
||||
import torch
|
||||
import torchvision
|
||||
import torchvision.transforms as transforms
|
||||
from PIL import Image
|
||||
from torchvision.datasets.folder import IMG_EXTENSIONS, pil_loader
|
||||
from torchvision.io import write_video
|
||||
from torchvision.utils import save_image
|
||||
|
||||
IMG_FPS = 120
|
||||
VID_EXTENSIONS = (".mp4", ".avi", ".mov", ".mkv")
|
||||
|
||||
regex = re.compile(
|
||||
r"^(?:http|ftp)s?://" # http:// or https://
|
||||
r"(?:(?:[A-Z0-9](?:[A-Z0-9-]{0,61}[A-Z0-9])?\.)+(?:[A-Z]{2,6}\.?|[A-Z0-9-]{2,}\.?)|" # domain...
|
||||
r"localhost|" # localhost...
|
||||
r"\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3})" # ...or ip
|
||||
r"(?::\d+)?" # optional port
|
||||
r"(?:/?|[/?]\S+)$",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
# H:W
|
||||
ASPECT_RATIO_MAP = {
|
||||
"3:8": "0.38",
|
||||
"9:21": "0.43",
|
||||
"12:25": "0.48",
|
||||
"1:2": "0.50",
|
||||
"9:17": "0.53",
|
||||
"27:50": "0.54",
|
||||
"9:16": "0.56",
|
||||
"5:8": "0.62",
|
||||
"2:3": "0.67",
|
||||
"3:4": "0.75",
|
||||
"1:1": "1.00",
|
||||
"4:3": "1.33",
|
||||
"3:2": "1.50",
|
||||
"16:9": "1.78",
|
||||
"17:9": "1.89",
|
||||
"2:1": "2.00",
|
||||
"50:27": "2.08",
|
||||
}
|
||||
|
||||
|
||||
# computed from above code
|
||||
# S = 8294400
|
||||
ASPECT_RATIO_4K = {
|
||||
"0.38": (1764, 4704),
|
||||
"0.43": (1886, 4400),
|
||||
"0.48": (1996, 4158),
|
||||
"0.50": (2036, 4072),
|
||||
"0.53": (2096, 3960),
|
||||
"0.54": (2118, 3918),
|
||||
"0.62": (2276, 3642),
|
||||
"0.56": (2160, 3840), # base
|
||||
"0.67": (2352, 3528),
|
||||
"0.75": (2494, 3326),
|
||||
"1.00": (2880, 2880),
|
||||
"1.33": (3326, 2494),
|
||||
"1.50": (3528, 2352),
|
||||
"1.78": (3840, 2160),
|
||||
"1.89": (3958, 2096),
|
||||
"2.00": (4072, 2036),
|
||||
"2.08": (4156, 1994),
|
||||
}
|
||||
|
||||
# S = 3686400
|
||||
ASPECT_RATIO_2K = {
|
||||
"0.38": (1176, 3136),
|
||||
"0.43": (1256, 2930),
|
||||
"0.48": (1330, 2770),
|
||||
"0.50": (1358, 2716),
|
||||
"0.53": (1398, 2640),
|
||||
"0.54": (1412, 2612),
|
||||
"0.56": (1440, 2560), # base
|
||||
"0.62": (1518, 2428),
|
||||
"0.67": (1568, 2352),
|
||||
"0.75": (1662, 2216),
|
||||
"1.00": (1920, 1920),
|
||||
"1.33": (2218, 1664),
|
||||
"1.50": (2352, 1568),
|
||||
"1.78": (2560, 1440),
|
||||
"1.89": (2638, 1396),
|
||||
"2.00": (2716, 1358),
|
||||
"2.08": (2772, 1330),
|
||||
}
|
||||
|
||||
# S = 2073600
|
||||
ASPECT_RATIO_1080P = {
|
||||
"0.38": (882, 2352),
|
||||
"0.43": (942, 2198),
|
||||
"0.48": (998, 2080),
|
||||
"0.50": (1018, 2036),
|
||||
"0.53": (1048, 1980),
|
||||
"0.54": (1058, 1958),
|
||||
"0.56": (1080, 1920), # base
|
||||
"0.62": (1138, 1820),
|
||||
"0.67": (1176, 1764),
|
||||
"0.75": (1248, 1664),
|
||||
"1.00": (1440, 1440),
|
||||
"1.33": (1662, 1246),
|
||||
"1.50": (1764, 1176),
|
||||
"1.78": (1920, 1080),
|
||||
"1.89": (1980, 1048),
|
||||
"2.00": (2036, 1018),
|
||||
"2.08": (2078, 998),
|
||||
}
|
||||
|
||||
# S = 921600
|
||||
ASPECT_RATIO_720P = {
|
||||
"0.38": (588, 1568),
|
||||
"0.43": (628, 1466),
|
||||
"0.48": (666, 1388),
|
||||
"0.50": (678, 1356),
|
||||
"0.53": (698, 1318),
|
||||
"0.54": (706, 1306),
|
||||
"0.56": (720, 1280), # base
|
||||
"0.62": (758, 1212),
|
||||
"0.67": (784, 1176),
|
||||
"0.75": (832, 1110),
|
||||
"1.00": (960, 960),
|
||||
"1.33": (1108, 832),
|
||||
"1.50": (1176, 784),
|
||||
"1.78": (1280, 720),
|
||||
"1.89": (1320, 698),
|
||||
"2.00": (1358, 680),
|
||||
"2.08": (1386, 666),
|
||||
}
|
||||
|
||||
# S = 409920
|
||||
ASPECT_RATIO_480P = {
|
||||
"0.38": (392, 1046),
|
||||
"0.43": (420, 980),
|
||||
"0.48": (444, 925),
|
||||
"0.50": (452, 904),
|
||||
"0.53": (466, 880),
|
||||
"0.54": (470, 870),
|
||||
"0.56": (480, 854), # base
|
||||
"0.62": (506, 810),
|
||||
"0.67": (522, 784),
|
||||
"0.75": (554, 738),
|
||||
"1.00": (640, 640),
|
||||
"1.33": (740, 555),
|
||||
"1.50": (784, 522),
|
||||
"1.78": (854, 480),
|
||||
"1.89": (880, 466),
|
||||
"2.00": (906, 454),
|
||||
"2.08": (924, 444),
|
||||
}
|
||||
|
||||
# S = 230400
|
||||
ASPECT_RATIO_360P = {
|
||||
"0.38": (294, 784),
|
||||
"0.43": (314, 732),
|
||||
"0.48": (332, 692),
|
||||
"0.50": (340, 680),
|
||||
"0.53": (350, 662),
|
||||
"0.54": (352, 652),
|
||||
"0.56": (360, 640), # base
|
||||
"0.62": (380, 608),
|
||||
"0.67": (392, 588),
|
||||
"0.75": (416, 554),
|
||||
"1.00": (480, 480),
|
||||
"1.33": (554, 416),
|
||||
"1.50": (588, 392),
|
||||
"1.78": (640, 360),
|
||||
"1.89": (660, 350),
|
||||
"2.00": (678, 340),
|
||||
"2.08": (692, 332),
|
||||
}
|
||||
|
||||
# S = 102240
|
||||
ASPECT_RATIO_240P = {
|
||||
"0.38": (196, 522),
|
||||
"0.43": (210, 490),
|
||||
"0.48": (222, 462),
|
||||
"0.50": (226, 452),
|
||||
"0.53": (232, 438),
|
||||
"0.54": (236, 436),
|
||||
"0.56": (240, 426), # base
|
||||
"0.62": (252, 404),
|
||||
"0.67": (262, 393),
|
||||
"0.75": (276, 368),
|
||||
"1.00": (320, 320),
|
||||
"1.33": (370, 278),
|
||||
"1.50": (392, 262),
|
||||
"1.78": (426, 240),
|
||||
"1.89": (440, 232),
|
||||
"2.00": (452, 226),
|
||||
"2.08": (462, 222),
|
||||
}
|
||||
|
||||
# S = 36864
|
||||
ASPECT_RATIO_144P = {
|
||||
"0.38": (117, 312),
|
||||
"0.43": (125, 291),
|
||||
"0.48": (133, 277),
|
||||
"0.50": (135, 270),
|
||||
"0.53": (139, 262),
|
||||
"0.54": (141, 260),
|
||||
"0.56": (144, 256), # base
|
||||
"0.62": (151, 241),
|
||||
"0.67": (156, 234),
|
||||
"0.75": (166, 221),
|
||||
"1.00": (192, 192),
|
||||
"1.33": (221, 165),
|
||||
"1.50": (235, 156),
|
||||
"1.78": (256, 144),
|
||||
"1.89": (263, 139),
|
||||
"2.00": (271, 135),
|
||||
"2.08": (277, 132),
|
||||
}
|
||||
|
||||
# from PixArt
|
||||
# S = 8294400
|
||||
ASPECT_RATIO_2880 = {
|
||||
"0.25": (1408, 5760),
|
||||
"0.26": (1408, 5568),
|
||||
"0.27": (1408, 5376),
|
||||
"0.28": (1408, 5184),
|
||||
"0.32": (1600, 4992),
|
||||
"0.33": (1600, 4800),
|
||||
"0.34": (1600, 4672),
|
||||
"0.4": (1792, 4480),
|
||||
"0.42": (1792, 4288),
|
||||
"0.47": (1920, 4096),
|
||||
"0.49": (1920, 3904),
|
||||
"0.51": (1920, 3776),
|
||||
"0.55": (2112, 3840),
|
||||
"0.59": (2112, 3584),
|
||||
"0.68": (2304, 3392),
|
||||
"0.72": (2304, 3200),
|
||||
"0.78": (2496, 3200),
|
||||
"0.83": (2496, 3008),
|
||||
"0.89": (2688, 3008),
|
||||
"0.93": (2688, 2880),
|
||||
"1.0": (2880, 2880),
|
||||
"1.07": (2880, 2688),
|
||||
"1.12": (3008, 2688),
|
||||
"1.21": (3008, 2496),
|
||||
"1.28": (3200, 2496),
|
||||
"1.39": (3200, 2304),
|
||||
"1.47": (3392, 2304),
|
||||
"1.7": (3584, 2112),
|
||||
"1.82": (3840, 2112),
|
||||
"2.03": (3904, 1920),
|
||||
"2.13": (4096, 1920),
|
||||
"2.39": (4288, 1792),
|
||||
"2.5": (4480, 1792),
|
||||
"2.92": (4672, 1600),
|
||||
"3.0": (4800, 1600),
|
||||
"3.12": (4992, 1600),
|
||||
"3.68": (5184, 1408),
|
||||
"3.82": (5376, 1408),
|
||||
"3.95": (5568, 1408),
|
||||
"4.0": (5760, 1408),
|
||||
}
|
||||
|
||||
# S = 4194304
|
||||
ASPECT_RATIO_2048 = {
|
||||
"0.25": (1024, 4096),
|
||||
"0.26": (1024, 3968),
|
||||
"0.27": (1024, 3840),
|
||||
"0.28": (1024, 3712),
|
||||
"0.32": (1152, 3584),
|
||||
"0.33": (1152, 3456),
|
||||
"0.35": (1152, 3328),
|
||||
"0.4": (1280, 3200),
|
||||
"0.42": (1280, 3072),
|
||||
"0.48": (1408, 2944),
|
||||
"0.5": (1408, 2816),
|
||||
"0.52": (1408, 2688),
|
||||
"0.57": (1536, 2688),
|
||||
"0.6": (1536, 2560),
|
||||
"0.68": (1664, 2432),
|
||||
"0.72": (1664, 2304),
|
||||
"0.78": (1792, 2304),
|
||||
"0.82": (1792, 2176),
|
||||
"0.88": (1920, 2176),
|
||||
"0.94": (1920, 2048),
|
||||
"1.0": (2048, 2048),
|
||||
"1.07": (2048, 1920),
|
||||
"1.13": (2176, 1920),
|
||||
"1.21": (2176, 1792),
|
||||
"1.29": (2304, 1792),
|
||||
"1.38": (2304, 1664),
|
||||
"1.46": (2432, 1664),
|
||||
"1.67": (2560, 1536),
|
||||
"1.75": (2688, 1536),
|
||||
"2.0": (2816, 1408),
|
||||
"2.09": (2944, 1408),
|
||||
"2.4": (3072, 1280),
|
||||
"2.5": (3200, 1280),
|
||||
"2.89": (3328, 1152),
|
||||
"3.0": (3456, 1152),
|
||||
"3.11": (3584, 1152),
|
||||
"3.62": (3712, 1024),
|
||||
"3.75": (3840, 1024),
|
||||
"3.88": (3968, 1024),
|
||||
"4.0": (4096, 1024),
|
||||
}
|
||||
|
||||
# S = 1048576
|
||||
ASPECT_RATIO_1024 = {
|
||||
"0.25": (512, 2048),
|
||||
"0.26": (512, 1984),
|
||||
"0.27": (512, 1920),
|
||||
"0.28": (512, 1856),
|
||||
"0.32": (576, 1792),
|
||||
"0.33": (576, 1728),
|
||||
"0.35": (576, 1664),
|
||||
"0.4": (640, 1600),
|
||||
"0.42": (640, 1536),
|
||||
"0.48": (704, 1472),
|
||||
"0.5": (704, 1408),
|
||||
"0.52": (704, 1344),
|
||||
"0.57": (768, 1344),
|
||||
"0.6": (768, 1280),
|
||||
"0.68": (832, 1216),
|
||||
"0.72": (832, 1152),
|
||||
"0.78": (896, 1152),
|
||||
"0.82": (896, 1088),
|
||||
"0.88": (960, 1088),
|
||||
"0.94": (960, 1024),
|
||||
"1.0": (1024, 1024),
|
||||
"1.07": (1024, 960),
|
||||
"1.13": (1088, 960),
|
||||
"1.21": (1088, 896),
|
||||
"1.29": (1152, 896),
|
||||
"1.38": (1152, 832),
|
||||
"1.46": (1216, 832),
|
||||
"1.67": (1280, 768),
|
||||
"1.75": (1344, 768),
|
||||
"2.0": (1408, 704),
|
||||
"2.09": (1472, 704),
|
||||
"2.4": (1536, 640),
|
||||
"2.5": (1600, 640),
|
||||
"2.89": (1664, 576),
|
||||
"3.0": (1728, 576),
|
||||
"3.11": (1792, 576),
|
||||
"3.62": (1856, 512),
|
||||
"3.75": (1920, 512),
|
||||
"3.88": (1984, 512),
|
||||
"4.0": (2048, 512),
|
||||
}
|
||||
|
||||
# S = 262144
|
||||
ASPECT_RATIO_512 = {
|
||||
"0.25": (256, 1024),
|
||||
"0.26": (256, 992),
|
||||
"0.27": (256, 960),
|
||||
"0.28": (256, 928),
|
||||
"0.32": (288, 896),
|
||||
"0.33": (288, 864),
|
||||
"0.35": (288, 832),
|
||||
"0.4": (320, 800),
|
||||
"0.42": (320, 768),
|
||||
"0.48": (352, 736),
|
||||
"0.5": (352, 704),
|
||||
"0.52": (352, 672),
|
||||
"0.57": (384, 672),
|
||||
"0.6": (384, 640),
|
||||
"0.68": (416, 608),
|
||||
"0.72": (416, 576),
|
||||
"0.78": (448, 576),
|
||||
"0.82": (448, 544),
|
||||
"0.88": (480, 544),
|
||||
"0.94": (480, 512),
|
||||
"1.0": (512, 512),
|
||||
"1.07": (512, 480),
|
||||
"1.13": (544, 480),
|
||||
"1.21": (544, 448),
|
||||
"1.29": (576, 448),
|
||||
"1.38": (576, 416),
|
||||
"1.46": (608, 416),
|
||||
"1.67": (640, 384),
|
||||
"1.75": (672, 384),
|
||||
"2.0": (704, 352),
|
||||
"2.09": (736, 352),
|
||||
"2.4": (768, 320),
|
||||
"2.5": (800, 320),
|
||||
"2.89": (832, 288),
|
||||
"3.0": (864, 288),
|
||||
"3.11": (896, 288),
|
||||
"3.62": (928, 256),
|
||||
"3.75": (960, 256),
|
||||
"3.88": (992, 256),
|
||||
"4.0": (1024, 256),
|
||||
}
|
||||
|
||||
# S = 65536
|
||||
ASPECT_RATIO_256 = {
|
||||
"0.25": (128, 512),
|
||||
"0.26": (128, 496),
|
||||
"0.27": (128, 480),
|
||||
"0.28": (128, 464),
|
||||
"0.32": (144, 448),
|
||||
"0.33": (144, 432),
|
||||
"0.35": (144, 416),
|
||||
"0.4": (160, 400),
|
||||
"0.42": (160, 384),
|
||||
"0.48": (176, 368),
|
||||
"0.5": (176, 352),
|
||||
"0.52": (176, 336),
|
||||
"0.57": (192, 336),
|
||||
"0.6": (192, 320),
|
||||
"0.68": (208, 304),
|
||||
"0.72": (208, 288),
|
||||
"0.78": (224, 288),
|
||||
"0.82": (224, 272),
|
||||
"0.88": (240, 272),
|
||||
"0.94": (240, 256),
|
||||
"1.0": (256, 256),
|
||||
"1.07": (256, 240),
|
||||
"1.13": (272, 240),
|
||||
"1.21": (272, 224),
|
||||
"1.29": (288, 224),
|
||||
"1.38": (288, 208),
|
||||
"1.46": (304, 208),
|
||||
"1.67": (320, 192),
|
||||
"1.75": (336, 192),
|
||||
"2.0": (352, 176),
|
||||
"2.09": (368, 176),
|
||||
"2.4": (384, 160),
|
||||
"2.5": (400, 160),
|
||||
"2.89": (416, 144),
|
||||
"3.0": (432, 144),
|
||||
"3.11": (448, 144),
|
||||
"3.62": (464, 128),
|
||||
"3.75": (480, 128),
|
||||
"3.88": (496, 128),
|
||||
"4.0": (512, 128),
|
||||
}
|
||||
|
||||
|
||||
def get_closest_ratio(height: float, width: float, ratios: dict):
|
||||
aspect_ratio = height / width
|
||||
closest_ratio = min(ratios.keys(), key=lambda ratio: abs(float(ratio) - aspect_ratio))
|
||||
return closest_ratio
|
||||
|
||||
|
||||
ASPECT_RATIOS = {
|
||||
"144p": (36864, ASPECT_RATIO_144P),
|
||||
"256": (65536, ASPECT_RATIO_256),
|
||||
"240p": (102240, ASPECT_RATIO_240P),
|
||||
"360p": (230400, ASPECT_RATIO_360P),
|
||||
"512": (262144, ASPECT_RATIO_512),
|
||||
"480p": (409920, ASPECT_RATIO_480P),
|
||||
"720p": (921600, ASPECT_RATIO_720P),
|
||||
"1024": (1048576, ASPECT_RATIO_1024),
|
||||
"1080p": (2073600, ASPECT_RATIO_1080P),
|
||||
"2k": (3686400, ASPECT_RATIO_2K),
|
||||
"2048": (4194304, ASPECT_RATIO_2048),
|
||||
"2880": (8294400, ASPECT_RATIO_2880),
|
||||
"4k": (8294400, ASPECT_RATIO_4K),
|
||||
}
|
||||
|
||||
|
||||
def get_image_size(resolution, ar_ratio):
|
||||
ar_key = ASPECT_RATIO_MAP[ar_ratio]
|
||||
rs_dict = ASPECT_RATIOS[resolution][1]
|
||||
assert ar_key in rs_dict, f"Aspect ratio {ar_ratio} not found for resolution {resolution}"
|
||||
return rs_dict[ar_key]
|
||||
|
||||
|
||||
NUM_FRAMES_MAP = {
|
||||
"1x": 51,
|
||||
"2x": 102,
|
||||
"4x": 204,
|
||||
"8x": 408,
|
||||
"16x": 816,
|
||||
"2s": 51,
|
||||
"4s": 102,
|
||||
"8s": 204,
|
||||
"16s": 408,
|
||||
"32s": 816,
|
||||
}
|
||||
|
||||
|
||||
def get_num_frames(num_frames):
|
||||
if num_frames in NUM_FRAMES_MAP:
|
||||
return NUM_FRAMES_MAP[num_frames]
|
||||
else:
|
||||
return int(num_frames)
|
||||
|
||||
|
||||
def save_sample(x, save_path=None, fps=8, normalize=True, value_range=(-1, 1), force_video=False, verbose=True):
|
||||
"""
|
||||
Args:
|
||||
x (Tensor): shape [C, T, H, W]
|
||||
"""
|
||||
assert x.ndim == 4
|
||||
|
||||
if not force_video and x.shape[1] == 1: # T = 1: save as image
|
||||
save_path += ".png"
|
||||
x = x.squeeze(1)
|
||||
save_image([x], save_path, normalize=normalize, value_range=value_range)
|
||||
else:
|
||||
save_path += ".mp4"
|
||||
if normalize:
|
||||
low, high = value_range
|
||||
x.clamp_(min=low, max=high)
|
||||
x.sub_(low).div_(max(high - low, 1e-5))
|
||||
|
||||
x = x.mul(255).add_(0.5).clamp_(0, 255).permute(1, 2, 3, 0).to("cpu", torch.uint8)
|
||||
write_video(save_path, x, fps=fps, video_codec="h264")
|
||||
if verbose:
|
||||
print(f"Saved to {save_path}")
|
||||
return save_path
|
||||
|
||||
|
||||
def is_url(url):
|
||||
return re.match(regex, url) is not None
|
||||
|
||||
|
||||
def download_url(input_path):
|
||||
output_dir = "cache"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
base_name = os.path.basename(input_path)
|
||||
output_path = os.path.join(output_dir, base_name)
|
||||
img_data = requests.get(input_path).content
|
||||
with open(output_path, "wb") as handler:
|
||||
handler.write(img_data)
|
||||
print(f"URL {input_path} downloaded to {output_path}")
|
||||
return output_path
|
||||
|
||||
|
||||
def get_transforms_video(name="center", image_size=(256, 256)):
|
||||
if name is None:
|
||||
return None
|
||||
elif name == "center":
|
||||
assert image_size[0] == image_size[1], "image_size must be square for center crop"
|
||||
transform_video = transforms.Compose(
|
||||
[
|
||||
ToTensorVideo(), # TCHW
|
||||
# video_transforms.RandomHorizontalFlipVideo(),
|
||||
UCFCenterCropVideo(image_size[0]),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
]
|
||||
)
|
||||
elif name == "resize_crop":
|
||||
transform_video = transforms.Compose(
|
||||
[
|
||||
ToTensorVideo(), # TCHW
|
||||
ResizeCrop(image_size),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
]
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Transform {name} not implemented")
|
||||
return transform_video
|
||||
|
||||
|
||||
def crop(clip, i, j, h, w):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
"""
|
||||
if len(clip.size()) != 4:
|
||||
raise ValueError("clip should be a 4D tensor")
|
||||
return clip[..., i : i + h, j : j + w]
|
||||
|
||||
|
||||
def center_crop(clip, crop_size):
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
th, tw = crop_size
|
||||
if h < th or w < tw:
|
||||
raise ValueError("height and width must be no smaller than crop_size")
|
||||
|
||||
i = int(round((h - th) / 2.0))
|
||||
j = int(round((w - tw) / 2.0))
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
def resize_scale(clip, target_size, interpolation_mode):
|
||||
if len(target_size) != 2:
|
||||
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
|
||||
H, W = clip.size(-2), clip.size(-1)
|
||||
scale_ = target_size[0] / min(H, W)
|
||||
return torch.nn.functional.interpolate(clip, scale_factor=scale_, mode=interpolation_mode, align_corners=False)
|
||||
|
||||
|
||||
class UCFCenterCropVideo:
|
||||
"""
|
||||
First scale to the specified size in equal proportion to the short edge,
|
||||
then center cropping
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: scale resized / center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
clip_resize = resize_scale(clip=clip, target_size=self.size, interpolation_mode=self.interpolation_mode)
|
||||
clip_center_crop = center_crop(clip_resize, self.size)
|
||||
return clip_center_crop
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
def _is_tensor_video_clip(clip):
|
||||
if not torch.is_tensor(clip):
|
||||
raise TypeError("clip should be Tensor. Got %s" % type(clip))
|
||||
|
||||
if not clip.ndimension() == 4:
|
||||
raise ValueError("clip should be 4D. Got %dD" % clip.dim())
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def to_tensor(clip):
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
permute the dimensions of clip tensor
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
_is_tensor_video_clip(clip)
|
||||
if not clip.dtype == torch.uint8:
|
||||
raise TypeError("clip tensor should have data type uint8. Got %s" % str(clip.dtype))
|
||||
# return clip.float().permute(3, 0, 1, 2) / 255.0
|
||||
return clip.float() / 255.0
|
||||
|
||||
|
||||
class ToTensorVideo:
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
permute the dimensions of clip tensor
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
return to_tensor(clip)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return self.__class__.__name__
|
||||
|
||||
|
||||
class ResizeCrop:
|
||||
def __init__(self, size):
|
||||
if isinstance(size, numbers.Number):
|
||||
self.size = (int(size), int(size))
|
||||
else:
|
||||
self.size = size
|
||||
|
||||
def __call__(self, clip):
|
||||
clip = resize_crop_to_fill(clip, self.size)
|
||||
return clip
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size})"
|
||||
|
||||
|
||||
def get_transforms_image(name="center", image_size=(256, 256)):
|
||||
if name is None:
|
||||
return None
|
||||
elif name == "center":
|
||||
assert image_size[0] == image_size[1], "Image size must be square for center crop"
|
||||
transform = transforms.Compose(
|
||||
[
|
||||
transforms.Lambda(lambda pil_image: center_crop_arr(pil_image, image_size[0])),
|
||||
# transforms.RandomHorizontalFlip(),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
]
|
||||
)
|
||||
elif name == "resize_crop":
|
||||
transform = transforms.Compose(
|
||||
[
|
||||
transforms.Lambda(lambda pil_image: resize_crop_to_fill(pil_image, image_size)),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
]
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Transform {name} not implemented")
|
||||
return transform
|
||||
|
||||
|
||||
def center_crop_arr(pil_image, image_size):
|
||||
"""
|
||||
Center cropping implementation from ADM.
|
||||
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
|
||||
"""
|
||||
while min(*pil_image.size) >= 2 * image_size:
|
||||
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size), resample=Image.BOX)
|
||||
|
||||
scale = image_size / min(*pil_image.size)
|
||||
pil_image = pil_image.resize(tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC)
|
||||
|
||||
arr = np.array(pil_image)
|
||||
crop_y = (arr.shape[0] - image_size) // 2
|
||||
crop_x = (arr.shape[1] - image_size) // 2
|
||||
return Image.fromarray(arr[crop_y : crop_y + image_size, crop_x : crop_x + image_size])
|
||||
|
||||
|
||||
def resize_crop_to_fill(pil_image, image_size):
|
||||
w, h = pil_image.size # PIL is (W, H)
|
||||
th, tw = image_size
|
||||
rh, rw = th / h, tw / w
|
||||
if rh > rw:
|
||||
sh, sw = th, round(w * rh)
|
||||
image = pil_image.resize((sw, sh), Image.BICUBIC)
|
||||
i = 0
|
||||
j = int(round((sw - tw) / 2.0))
|
||||
else:
|
||||
sh, sw = round(h * rw), tw
|
||||
image = pil_image.resize((sw, sh), Image.BICUBIC)
|
||||
i = int(round((sh - th) / 2.0))
|
||||
j = 0
|
||||
arr = np.array(image)
|
||||
assert i + th <= arr.shape[0] and j + tw <= arr.shape[1]
|
||||
return Image.fromarray(arr[i : i + th, j : j + tw])
|
||||
|
||||
|
||||
def read_video_from_path(path, transform=None, transform_name="center", image_size=(256, 256)):
|
||||
vframes, aframes, info = torchvision.io.read_video(filename=path, pts_unit="sec", output_format="TCHW")
|
||||
if transform is None:
|
||||
transform = get_transforms_video(image_size=image_size, name=transform_name)
|
||||
video = transform(vframes) # T C H W
|
||||
video = video.permute(1, 0, 2, 3)
|
||||
return video
|
||||
|
||||
|
||||
def read_from_path(path, image_size, transform_name="center"):
|
||||
if is_url(path):
|
||||
path = download_url(path)
|
||||
ext = os.path.splitext(path)[-1].lower()
|
||||
if ext.lower() in VID_EXTENSIONS:
|
||||
return read_video_from_path(path, image_size=image_size, transform_name=transform_name)
|
||||
else:
|
||||
assert ext.lower() in IMG_EXTENSIONS, f"Unsupported file format: {ext}"
|
||||
return read_image_from_path(path, image_size=image_size, transform_name=transform_name)
|
||||
|
||||
|
||||
def read_image_from_path(path, transform=None, transform_name="center", num_frames=1, image_size=(256, 256)):
|
||||
image = pil_loader(path)
|
||||
if transform is None:
|
||||
transform = get_transforms_image(image_size=image_size, name=transform_name)
|
||||
image = transform(image)
|
||||
video = image.unsqueeze(0).repeat(num_frames, 1, 1, 1)
|
||||
video = video.permute(1, 0, 2, 3)
|
||||
return video
|
||||
Executable
+585
@@ -0,0 +1,585 @@
|
||||
# Adapted from OpenSora and DiT
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# DiT: https://github.com/facebookresearch/DiT
|
||||
# OpenSora: https://github.com/hpcaitech/Open-Sora
|
||||
# --------------------------------------------------------
|
||||
|
||||
import html
|
||||
import math
|
||||
import re
|
||||
|
||||
import ftfy
|
||||
import numpy
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import transformers
|
||||
from timm.models.vision_transformer import Mlp
|
||||
from transformers import AutoTokenizer, CLIPTextModel, CLIPTokenizer, T5EncoderModel
|
||||
|
||||
from opendit.modules.embed import get_1d_sincos_pos_embed_from_grid, get_2d_sincos_pos_embed_from_grid
|
||||
|
||||
transformers.logging.set_verbosity_error()
|
||||
|
||||
|
||||
# ===============================================
|
||||
# Text Embed
|
||||
# ===============================================
|
||||
|
||||
|
||||
class AbstractEncoder(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def encode(self, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class FrozenCLIPEmbedder(AbstractEncoder):
|
||||
"""Uses the CLIP transformer encoder for text (from Hugging Face)"""
|
||||
|
||||
def __init__(self, path="openai/clip-vit-huge-patch14", device="cuda", max_length=77):
|
||||
super().__init__()
|
||||
self.tokenizer = CLIPTokenizer.from_pretrained(path)
|
||||
self.transformer = CLIPTextModel.from_pretrained(path)
|
||||
self.device = device
|
||||
self.max_length = max_length
|
||||
self._freeze()
|
||||
|
||||
def _freeze(self):
|
||||
self.transformer = self.transformer.eval()
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, text):
|
||||
batch_encoding = self.tokenizer(
|
||||
text,
|
||||
truncation=True,
|
||||
max_length=self.max_length,
|
||||
return_length=True,
|
||||
return_overflowing_tokens=False,
|
||||
padding="max_length",
|
||||
return_tensors="pt",
|
||||
)
|
||||
tokens = batch_encoding["input_ids"].to(self.device)
|
||||
outputs = self.transformer(input_ids=tokens)
|
||||
|
||||
z = outputs.last_hidden_state
|
||||
pooled_z = outputs.pooler_output
|
||||
return z, pooled_z
|
||||
|
||||
def encode(self, text):
|
||||
return self(text)
|
||||
|
||||
|
||||
class TextEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds text prompt into vector representations. Also handles text dropout for classifier-free guidance.
|
||||
"""
|
||||
|
||||
def __init__(self, path, hidden_size, dropout_prob=0.1):
|
||||
super().__init__()
|
||||
self.text_encoder = FrozenCLIPEmbedder(path=path)
|
||||
self.dropout_prob = dropout_prob
|
||||
|
||||
output_dim = self.text_encoder.transformer.config.hidden_size
|
||||
self.output_projection = nn.Linear(output_dim, hidden_size)
|
||||
|
||||
def token_drop(self, text_prompts, force_drop_ids=None):
|
||||
"""
|
||||
Drops text to enable classifier-free guidance.
|
||||
"""
|
||||
if force_drop_ids is None:
|
||||
drop_ids = numpy.random.uniform(0, 1, len(text_prompts)) < self.dropout_prob
|
||||
else:
|
||||
# TODO
|
||||
drop_ids = force_drop_ids == 1
|
||||
labels = list(numpy.where(drop_ids, "", text_prompts))
|
||||
# print(labels)
|
||||
return labels
|
||||
|
||||
def forward(self, text_prompts, train, force_drop_ids=None):
|
||||
use_dropout = self.dropout_prob > 0
|
||||
if (train and use_dropout) or (force_drop_ids is not None):
|
||||
text_prompts = self.token_drop(text_prompts, force_drop_ids)
|
||||
embeddings, pooled_embeddings = self.text_encoder(text_prompts)
|
||||
# return embeddings, pooled_embeddings
|
||||
text_embeddings = self.output_projection(pooled_embeddings)
|
||||
return text_embeddings
|
||||
|
||||
|
||||
class CaptionEmbedder(nn.Module):
|
||||
"""
|
||||
copied from https://github.com/hpcaitech/Open-Sora
|
||||
|
||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
||||
"""
|
||||
|
||||
def __init__(self, in_channels, hidden_size, uncond_prob, act_layer=nn.GELU(approximate="tanh"), token_num=120):
|
||||
super().__init__()
|
||||
|
||||
self.y_proj = Mlp(
|
||||
in_features=in_channels, hidden_features=hidden_size, out_features=hidden_size, act_layer=act_layer, drop=0
|
||||
)
|
||||
self.register_buffer("y_embedding", nn.Parameter(torch.randn(token_num, in_channels) / in_channels**0.5))
|
||||
self.uncond_prob = uncond_prob
|
||||
|
||||
def token_drop(self, caption, force_drop_ids=None):
|
||||
"""
|
||||
Drops labels to enable classifier-free guidance.
|
||||
"""
|
||||
if force_drop_ids is None:
|
||||
drop_ids = torch.rand(caption.shape[0]).cuda() < self.uncond_prob
|
||||
else:
|
||||
drop_ids = force_drop_ids == 1
|
||||
caption = torch.where(drop_ids[:, None, None, None], self.y_embedding, caption)
|
||||
return caption
|
||||
|
||||
def forward(self, caption, train, force_drop_ids=None):
|
||||
if train:
|
||||
assert caption.shape[2:] == self.y_embedding.shape
|
||||
use_dropout = self.uncond_prob > 0
|
||||
if (train and use_dropout) or (force_drop_ids is not None):
|
||||
caption = self.token_drop(caption, force_drop_ids)
|
||||
caption = self.y_proj(caption)
|
||||
return caption
|
||||
|
||||
|
||||
class T5Embedder:
|
||||
available_models = ["DeepFloyd/t5-v1_1-xxl"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
device,
|
||||
from_pretrained=None,
|
||||
*,
|
||||
cache_dir=None,
|
||||
hf_token=None,
|
||||
use_text_preprocessing=True,
|
||||
t5_model_kwargs=None,
|
||||
torch_dtype=None,
|
||||
use_offload_folder=None,
|
||||
model_max_length=120,
|
||||
local_files_only=False,
|
||||
):
|
||||
self.device = torch.device(device)
|
||||
self.torch_dtype = torch_dtype or torch.bfloat16
|
||||
self.cache_dir = cache_dir
|
||||
|
||||
if t5_model_kwargs is None:
|
||||
t5_model_kwargs = {
|
||||
"low_cpu_mem_usage": True,
|
||||
"torch_dtype": self.torch_dtype,
|
||||
}
|
||||
|
||||
if use_offload_folder is not None:
|
||||
t5_model_kwargs["offload_folder"] = use_offload_folder
|
||||
t5_model_kwargs["device_map"] = {
|
||||
"shared": self.device,
|
||||
"encoder.embed_tokens": self.device,
|
||||
"encoder.block.0": self.device,
|
||||
"encoder.block.1": self.device,
|
||||
"encoder.block.2": self.device,
|
||||
"encoder.block.3": self.device,
|
||||
"encoder.block.4": self.device,
|
||||
"encoder.block.5": self.device,
|
||||
"encoder.block.6": self.device,
|
||||
"encoder.block.7": self.device,
|
||||
"encoder.block.8": self.device,
|
||||
"encoder.block.9": self.device,
|
||||
"encoder.block.10": self.device,
|
||||
"encoder.block.11": self.device,
|
||||
"encoder.block.12": "disk",
|
||||
"encoder.block.13": "disk",
|
||||
"encoder.block.14": "disk",
|
||||
"encoder.block.15": "disk",
|
||||
"encoder.block.16": "disk",
|
||||
"encoder.block.17": "disk",
|
||||
"encoder.block.18": "disk",
|
||||
"encoder.block.19": "disk",
|
||||
"encoder.block.20": "disk",
|
||||
"encoder.block.21": "disk",
|
||||
"encoder.block.22": "disk",
|
||||
"encoder.block.23": "disk",
|
||||
"encoder.final_layer_norm": "disk",
|
||||
"encoder.dropout": "disk",
|
||||
}
|
||||
else:
|
||||
t5_model_kwargs["device_map"] = {
|
||||
"shared": self.device,
|
||||
"encoder": self.device,
|
||||
}
|
||||
|
||||
self.use_text_preprocessing = use_text_preprocessing
|
||||
self.hf_token = hf_token
|
||||
|
||||
assert from_pretrained in self.available_models
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(
|
||||
from_pretrained,
|
||||
cache_dir=cache_dir,
|
||||
local_files_only=local_files_only,
|
||||
)
|
||||
self.model = T5EncoderModel.from_pretrained(
|
||||
from_pretrained,
|
||||
cache_dir=cache_dir,
|
||||
local_files_only=local_files_only,
|
||||
**t5_model_kwargs,
|
||||
).eval()
|
||||
self.model_max_length = model_max_length
|
||||
|
||||
def get_text_embeddings(self, texts):
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
texts,
|
||||
max_length=self.model_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
input_ids = text_tokens_and_mask["input_ids"].to(self.device)
|
||||
attention_mask = text_tokens_and_mask["attention_mask"].to(self.device)
|
||||
with torch.no_grad():
|
||||
text_encoder_embs = self.model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
)["last_hidden_state"].detach()
|
||||
return text_encoder_embs, attention_mask
|
||||
|
||||
|
||||
class T5Encoder:
|
||||
def __init__(
|
||||
self,
|
||||
from_pretrained="DeepFloyd/t5-v1_1-xxl",
|
||||
model_max_length=120,
|
||||
device="cuda",
|
||||
dtype=torch.float,
|
||||
shardformer=False,
|
||||
):
|
||||
assert from_pretrained is not None, "Please specify the path to the T5 model"
|
||||
|
||||
self.t5 = T5Embedder(
|
||||
device=device,
|
||||
torch_dtype=dtype,
|
||||
from_pretrained=from_pretrained,
|
||||
model_max_length=model_max_length,
|
||||
)
|
||||
self.t5.model.to(dtype=dtype)
|
||||
self.y_embedder = None
|
||||
|
||||
self.model_max_length = model_max_length
|
||||
self.output_dim = self.t5.model.config.d_model
|
||||
|
||||
if shardformer:
|
||||
self.shardformer_t5()
|
||||
|
||||
def shardformer_t5(self):
|
||||
from colossalai.shardformer import ShardConfig, ShardFormer
|
||||
|
||||
from opendit.core.shardformer.t5.policy import T5EncoderPolicy
|
||||
from opendit.utils.utils import requires_grad
|
||||
|
||||
shard_config = ShardConfig(
|
||||
tensor_parallel_process_group=None,
|
||||
pipeline_stage_manager=None,
|
||||
enable_tensor_parallelism=False,
|
||||
enable_fused_normalization=False,
|
||||
enable_flash_attention=False,
|
||||
enable_jit_fused=True,
|
||||
enable_sequence_parallelism=False,
|
||||
enable_sequence_overlap=False,
|
||||
)
|
||||
shard_former = ShardFormer(shard_config=shard_config)
|
||||
optim_model, _ = shard_former.optimize(self.t5.model, policy=T5EncoderPolicy())
|
||||
self.t5.model = optim_model.half()
|
||||
|
||||
# ensure the weights are frozen
|
||||
requires_grad(self.t5.model, False)
|
||||
|
||||
def encode(self, text):
|
||||
caption_embs, emb_masks = self.t5.get_text_embeddings(text)
|
||||
caption_embs = caption_embs[:, None]
|
||||
return dict(y=caption_embs, mask=emb_masks)
|
||||
|
||||
def null(self, n):
|
||||
null_y = self.y_embedder.y_embedding[None].repeat(n, 1, 1)[:, None]
|
||||
return null_y
|
||||
|
||||
|
||||
def basic_clean(text):
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
|
||||
BAD_PUNCT_REGEX = re.compile(
|
||||
r"[" + "#®•©™&@·º½¾¿¡§~" + "\)" + "\(" + "\]" + "\[" + "\}" + "\{" + "\|" + "\\" + "\/" + "\*" + r"]{1,}"
|
||||
) # noqa
|
||||
|
||||
|
||||
def clean_caption(caption):
|
||||
import urllib.parse as ul
|
||||
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
caption = str(caption)
|
||||
caption = ul.unquote_plus(caption)
|
||||
caption = caption.strip().lower()
|
||||
caption = re.sub("<person>", "person", caption)
|
||||
# urls:
|
||||
caption = re.sub(
|
||||
r"\b((?:https?:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))", # noqa
|
||||
"",
|
||||
caption,
|
||||
) # regex for urls
|
||||
caption = re.sub(
|
||||
r"\b((?:www:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))", # noqa
|
||||
"",
|
||||
caption,
|
||||
) # regex for urls
|
||||
# html:
|
||||
caption = BeautifulSoup(caption, features="html.parser").text
|
||||
|
||||
# @<nickname>
|
||||
caption = re.sub(r"@[\w\d]+\b", "", caption)
|
||||
|
||||
# 31C0—31EF CJK Strokes
|
||||
# 31F0—31FF Katakana Phonetic Extensions
|
||||
# 3200—32FF Enclosed CJK Letters and Months
|
||||
# 3300—33FF CJK Compatibility
|
||||
# 3400—4DBF CJK Unified Ideographs Extension A
|
||||
# 4DC0—4DFF Yijing Hexagram Symbols
|
||||
# 4E00—9FFF CJK Unified Ideographs
|
||||
caption = re.sub(r"[\u31c0-\u31ef]+", "", caption)
|
||||
caption = re.sub(r"[\u31f0-\u31ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3200-\u32ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3300-\u33ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3400-\u4dbf]+", "", caption)
|
||||
caption = re.sub(r"[\u4dc0-\u4dff]+", "", caption)
|
||||
caption = re.sub(r"[\u4e00-\u9fff]+", "", caption)
|
||||
#######################################################
|
||||
|
||||
# все виды тире / all types of dash --> "-"
|
||||
caption = re.sub(
|
||||
r"[\u002D\u058A\u05BE\u1400\u1806\u2010-\u2015\u2E17\u2E1A\u2E3A\u2E3B\u2E40\u301C\u3030\u30A0\uFE31\uFE32\uFE58\uFE63\uFF0D]+", # noqa
|
||||
"-",
|
||||
caption,
|
||||
)
|
||||
|
||||
# кавычки к одному стандарту
|
||||
caption = re.sub(r"[`´«»“”¨]", '"', caption)
|
||||
caption = re.sub(r"[‘’]", "'", caption)
|
||||
|
||||
# "
|
||||
caption = re.sub(r""?", "", caption)
|
||||
# &
|
||||
caption = re.sub(r"&", "", caption)
|
||||
|
||||
# ip adresses:
|
||||
caption = re.sub(r"\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}", " ", caption)
|
||||
|
||||
# article ids:
|
||||
caption = re.sub(r"\d:\d\d\s+$", "", caption)
|
||||
|
||||
# \n
|
||||
caption = re.sub(r"\\n", " ", caption)
|
||||
|
||||
# "#123"
|
||||
caption = re.sub(r"#\d{1,3}\b", "", caption)
|
||||
# "#12345.."
|
||||
caption = re.sub(r"#\d{5,}\b", "", caption)
|
||||
# "123456.."
|
||||
caption = re.sub(r"\b\d{6,}\b", "", caption)
|
||||
# filenames:
|
||||
caption = re.sub(r"[\S]+\.(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)", "", caption)
|
||||
|
||||
#
|
||||
caption = re.sub(r"[\"\']{2,}", r'"', caption) # """AUSVERKAUFT"""
|
||||
caption = re.sub(r"[\.]{2,}", r" ", caption) # """AUSVERKAUFT"""
|
||||
|
||||
caption = re.sub(BAD_PUNCT_REGEX, r" ", caption) # ***AUSVERKAUFT***, #AUSVERKAUFT
|
||||
caption = re.sub(r"\s+\.\s+", r" ", caption) # " . "
|
||||
|
||||
# this-is-my-cute-cat / this_is_my_cute_cat
|
||||
regex2 = re.compile(r"(?:\-|\_)")
|
||||
if len(re.findall(regex2, caption)) > 3:
|
||||
caption = re.sub(regex2, " ", caption)
|
||||
|
||||
caption = basic_clean(caption)
|
||||
|
||||
caption = re.sub(r"\b[a-zA-Z]{1,3}\d{3,15}\b", "", caption) # jc6640
|
||||
caption = re.sub(r"\b[a-zA-Z]+\d+[a-zA-Z]+\b", "", caption) # jc6640vc
|
||||
caption = re.sub(r"\b\d+[a-zA-Z]+\d+\b", "", caption) # 6640vc231
|
||||
|
||||
caption = re.sub(r"(worldwide\s+)?(free\s+)?shipping", "", caption)
|
||||
caption = re.sub(r"(free\s)?download(\sfree)?", "", caption)
|
||||
caption = re.sub(r"\bclick\b\s(?:for|on)\s\w+", "", caption)
|
||||
caption = re.sub(r"\b(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)(\simage[s]?)?", "", caption)
|
||||
caption = re.sub(r"\bpage\s+\d+\b", "", caption)
|
||||
|
||||
caption = re.sub(r"\b\d*[a-zA-Z]+\d+[a-zA-Z]+\d+[a-zA-Z\d]*\b", r" ", caption) # j2d1a2a...
|
||||
|
||||
caption = re.sub(r"\b\d+\.?\d*[xх×]\d+\.?\d*\b", "", caption)
|
||||
|
||||
caption = re.sub(r"\b\s+\:\s+", r": ", caption)
|
||||
caption = re.sub(r"(\D[,\./])\b", r"\1 ", caption)
|
||||
caption = re.sub(r"\s+", " ", caption)
|
||||
|
||||
caption.strip()
|
||||
|
||||
caption = re.sub(r"^[\"\']([\w\W]+)[\"\']$", r"\1", caption)
|
||||
caption = re.sub(r"^[\'\_,\-\:;]", r"", caption)
|
||||
caption = re.sub(r"[\'\_,\-\:\-\+]$", r"", caption)
|
||||
caption = re.sub(r"^\.\S+$", "", caption)
|
||||
|
||||
return caption.strip()
|
||||
|
||||
|
||||
def text_preprocessing(text, use_text_preprocessing: bool = True):
|
||||
if use_text_preprocessing:
|
||||
# The exact text cleaning as was in the training stage:
|
||||
text = clean_caption(text)
|
||||
text = clean_caption(text)
|
||||
return text
|
||||
else:
|
||||
return text.lower().strip()
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__()
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
||||
)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
|
||||
@staticmethod
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param t: 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, D) Tensor of positional embeddings.
|
||||
"""
|
||||
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
||||
half = dim // 2
|
||||
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half)
|
||||
freqs = freqs.to(device=t.device)
|
||||
args = t[:, 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)
|
||||
return embedding
|
||||
|
||||
def forward(self, t, dtype):
|
||||
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
||||
if t_freq.dtype != dtype:
|
||||
t_freq = t_freq.to(dtype)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
|
||||
|
||||
# ===============================================
|
||||
# Sine/Cosine Positional Embedding Functions
|
||||
# ===============================================
|
||||
|
||||
|
||||
def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0, scale=1.0, base_size=None):
|
||||
"""
|
||||
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)
|
||||
"""
|
||||
if not isinstance(grid_size, tuple):
|
||||
grid_size = (grid_size, grid_size)
|
||||
|
||||
grid_h = np.arange(grid_size[0], dtype=np.float32) / scale
|
||||
grid_w = np.arange(grid_size[1], dtype=np.float32) / scale
|
||||
if base_size is not None:
|
||||
grid_h *= base_size / grid_size[0]
|
||||
grid_w *= base_size / grid_size[1]
|
||||
grid = np.meshgrid(grid_w, grid_h) # here w goes first
|
||||
grid = np.stack(grid, axis=0)
|
||||
|
||||
grid = grid.reshape([2, 1, grid_size[1], grid_size[0]])
|
||||
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
|
||||
if cls_token and extra_tokens > 0:
|
||||
pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0)
|
||||
return pos_embed
|
||||
|
||||
|
||||
def get_1d_sincos_pos_embed(embed_dim, length, scale=1.0):
|
||||
pos = np.arange(0, length)[..., None] / scale
|
||||
return get_1d_sincos_pos_embed_from_grid(embed_dim, pos)
|
||||
|
||||
|
||||
# ===============================================
|
||||
# Patch Embed
|
||||
# ===============================================
|
||||
|
||||
|
||||
class PatchEmbed3D(nn.Module):
|
||||
"""Video to Patch Embedding.
|
||||
|
||||
Args:
|
||||
patch_size (int): Patch token size. Default: (2,4,4).
|
||||
in_chans (int): Number of input video channels. Default: 3.
|
||||
embed_dim (int): Number of linear projection output channels. Default: 96.
|
||||
norm_layer (nn.Module, optional): Normalization layer. Default: None
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=(2, 4, 4),
|
||||
in_chans=3,
|
||||
embed_dim=96,
|
||||
norm_layer=None,
|
||||
flatten=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.flatten = flatten
|
||||
|
||||
self.in_chans = in_chans
|
||||
self.embed_dim = embed_dim
|
||||
|
||||
self.proj = nn.Conv3d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
|
||||
if norm_layer is not None:
|
||||
self.norm = norm_layer(embed_dim)
|
||||
else:
|
||||
self.norm = None
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward function."""
|
||||
# padding
|
||||
_, _, D, H, W = x.size()
|
||||
if W % self.patch_size[2] != 0:
|
||||
x = F.pad(x, (0, self.patch_size[2] - W % self.patch_size[2]))
|
||||
if H % self.patch_size[1] != 0:
|
||||
x = F.pad(x, (0, 0, 0, self.patch_size[1] - H % self.patch_size[1]))
|
||||
if D % self.patch_size[0] != 0:
|
||||
x = F.pad(x, (0, 0, 0, 0, 0, self.patch_size[0] - D % self.patch_size[0]))
|
||||
|
||||
x = self.proj(x) # (B C T H W)
|
||||
if self.norm is not None:
|
||||
D, Wh, Ww = x.size(2), x.size(3), x.size(4)
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
x = self.norm(x)
|
||||
x = x.transpose(1, 2).view(-1, self.embed_dim, D, Wh, Ww)
|
||||
if self.flatten:
|
||||
x = x.flatten(2).transpose(1, 2) # BCTHW -> BNC
|
||||
return x
|
||||
Executable
+348
@@ -0,0 +1,348 @@
|
||||
# Adapted from OpenSora
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# OpenSora: https://github.com/hpcaitech/Open-Sora
|
||||
# --------------------------------------------------------
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
|
||||
import torch
|
||||
|
||||
from .datasets import IMG_FPS, read_from_path
|
||||
|
||||
|
||||
def prepare_multi_resolution_info(info_type, batch_size, image_size, num_frames, fps, device, dtype):
|
||||
if info_type is None:
|
||||
return dict()
|
||||
elif info_type == "PixArtMS":
|
||||
hw = torch.tensor([image_size], device=device, dtype=dtype).repeat(batch_size, 1)
|
||||
ar = torch.tensor([[image_size[0] / image_size[1]]], device=device, dtype=dtype).repeat(batch_size, 1)
|
||||
return dict(ar=ar, hw=hw)
|
||||
elif info_type in ["STDiT2", "OpenSora"]:
|
||||
fps = fps if num_frames > 1 else IMG_FPS
|
||||
fps = torch.tensor([fps], device=device, dtype=dtype).repeat(batch_size)
|
||||
height = torch.tensor([image_size[0]], device=device, dtype=dtype).repeat(batch_size)
|
||||
width = torch.tensor([image_size[1]], device=device, dtype=dtype).repeat(batch_size)
|
||||
num_frames = torch.tensor([num_frames], device=device, dtype=dtype).repeat(batch_size)
|
||||
ar = torch.tensor([image_size[0] / image_size[1]], device=device, dtype=dtype).repeat(batch_size)
|
||||
return dict(height=height, width=width, num_frames=num_frames, ar=ar, fps=fps)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def load_prompts(prompt_path, start_idx=None, end_idx=None):
|
||||
with open(prompt_path, "r") as f:
|
||||
prompts = [line.strip() for line in f.readlines()]
|
||||
prompts = prompts[start_idx:end_idx]
|
||||
return prompts
|
||||
|
||||
|
||||
def get_save_path_name(
|
||||
save_dir,
|
||||
sample_name=None, # prefix
|
||||
sample_idx=None, # sample index
|
||||
prompt=None, # used prompt
|
||||
prompt_as_path=False, # use prompt as path
|
||||
num_sample=1, # number of samples to generate for one prompt
|
||||
k=None, # kth sample
|
||||
):
|
||||
if sample_name is None:
|
||||
sample_name = "" if prompt_as_path else "sample"
|
||||
sample_name_suffix = prompt if prompt_as_path else f"_{sample_idx:04d}"
|
||||
save_path = os.path.join(save_dir, f"{sample_name}{sample_name_suffix[:40]}")
|
||||
if num_sample != 1:
|
||||
save_path = f"{save_path}-{k}"
|
||||
return save_path
|
||||
|
||||
|
||||
def get_eval_save_path_name(
|
||||
save_dir,
|
||||
id, # add id parameter
|
||||
sample_name=None, # prefix
|
||||
sample_idx=None, # sample index
|
||||
prompt=None, # used prompt
|
||||
prompt_as_path=False, # use prompt as path
|
||||
num_sample=1, # number of samples to generate for one prompt
|
||||
k=None, # kth sample
|
||||
):
|
||||
if sample_name is None:
|
||||
sample_name = "" if prompt_as_path else "sample"
|
||||
save_path = os.path.join(save_dir, f"{id}")
|
||||
if num_sample != 1:
|
||||
save_path = f"{save_path}-{k}"
|
||||
return save_path
|
||||
|
||||
|
||||
def append_score_to_prompts(prompts, aes=None, flow=None, camera_motion=None):
|
||||
new_prompts = []
|
||||
for prompt in prompts:
|
||||
new_prompt = prompt
|
||||
if aes is not None and "aesthetic score:" not in prompt:
|
||||
new_prompt = f"{new_prompt} aesthetic score: {aes:.1f}."
|
||||
if flow is not None and "motion score:" not in prompt:
|
||||
new_prompt = f"{new_prompt} motion score: {flow:.1f}."
|
||||
if camera_motion is not None and "camera motion:" not in prompt:
|
||||
new_prompt = f"{new_prompt} camera motion: {camera_motion}."
|
||||
new_prompts.append(new_prompt)
|
||||
return new_prompts
|
||||
|
||||
|
||||
def extract_json_from_prompts(prompts, reference, mask_strategy):
|
||||
ret_prompts = []
|
||||
for i, prompt in enumerate(prompts):
|
||||
parts = re.split(r"(?=[{])", prompt)
|
||||
assert len(parts) <= 2, f"Invalid prompt: {prompt}"
|
||||
ret_prompts.append(parts[0])
|
||||
if len(parts) > 1:
|
||||
additional_info = json.loads(parts[1])
|
||||
for key in additional_info:
|
||||
assert key in ["reference_path", "mask_strategy"], f"Invalid key: {key}"
|
||||
if key == "reference_path":
|
||||
reference[i] = additional_info[key]
|
||||
elif key == "mask_strategy":
|
||||
mask_strategy[i] = additional_info[key]
|
||||
return ret_prompts, reference, mask_strategy
|
||||
|
||||
|
||||
def collect_references_batch(reference_paths, vae, image_size):
|
||||
refs_x = [] # refs_x: [batch, ref_num, C, T, H, W]
|
||||
for reference_path in reference_paths:
|
||||
if reference_path == "":
|
||||
refs_x.append([])
|
||||
continue
|
||||
ref_path = reference_path.split(";")
|
||||
ref = []
|
||||
for r_path in ref_path:
|
||||
r = read_from_path(r_path, image_size, transform_name="resize_crop")
|
||||
r_x = vae.encode(r.unsqueeze(0).to(vae.device, vae.dtype))
|
||||
r_x = r_x.squeeze(0)
|
||||
ref.append(r_x)
|
||||
refs_x.append(ref)
|
||||
return refs_x
|
||||
|
||||
|
||||
def extract_prompts_loop(prompts, num_loop):
|
||||
ret_prompts = []
|
||||
for prompt in prompts:
|
||||
if prompt.startswith("|0|"):
|
||||
prompt_list = prompt.split("|")[1:]
|
||||
text_list = []
|
||||
for i in range(0, len(prompt_list), 2):
|
||||
start_loop = int(prompt_list[i])
|
||||
text = prompt_list[i + 1]
|
||||
end_loop = int(prompt_list[i + 2]) if i + 2 < len(prompt_list) else num_loop + 1
|
||||
text_list.extend([text] * (end_loop - start_loop))
|
||||
prompt = text_list[num_loop]
|
||||
ret_prompts.append(prompt)
|
||||
return ret_prompts
|
||||
|
||||
|
||||
def split_prompt(prompt_text):
|
||||
if prompt_text.startswith("|0|"):
|
||||
# this is for prompts which look like
|
||||
# |0| a beautiful day |1| a sunny day |2| a rainy day
|
||||
# we want to parse it into a list of prompts with the loop index
|
||||
prompt_list = prompt_text.split("|")[1:]
|
||||
text_list = []
|
||||
loop_idx = []
|
||||
for i in range(0, len(prompt_list), 2):
|
||||
start_loop = int(prompt_list[i])
|
||||
text = prompt_list[i + 1].strip()
|
||||
text_list.append(text)
|
||||
loop_idx.append(start_loop)
|
||||
return text_list, loop_idx
|
||||
else:
|
||||
return [prompt_text], None
|
||||
|
||||
|
||||
def merge_prompt(text_list, loop_idx_list=None):
|
||||
if loop_idx_list is None:
|
||||
return text_list[0]
|
||||
else:
|
||||
prompt = ""
|
||||
for i, text in enumerate(text_list):
|
||||
prompt += f"|{loop_idx_list[i]}|{text}"
|
||||
return prompt
|
||||
|
||||
|
||||
MASK_DEFAULT = ["0", "0", "0", "0", "1", "0"]
|
||||
|
||||
|
||||
def parse_mask_strategy(mask_strategy):
|
||||
mask_batch = []
|
||||
if mask_strategy == "" or mask_strategy is None:
|
||||
return mask_batch
|
||||
|
||||
mask_strategy = mask_strategy.split(";")
|
||||
for mask in mask_strategy:
|
||||
mask_group = mask.split(",")
|
||||
num_group = len(mask_group)
|
||||
assert num_group >= 1 and num_group <= 6, f"Invalid mask strategy: {mask}"
|
||||
mask_group.extend(MASK_DEFAULT[num_group:])
|
||||
for i in range(5):
|
||||
mask_group[i] = int(mask_group[i])
|
||||
mask_group[5] = float(mask_group[5])
|
||||
mask_batch.append(mask_group)
|
||||
return mask_batch
|
||||
|
||||
|
||||
def find_nearest_point(value, point, max_value):
|
||||
t = value // point
|
||||
if value % point > point / 2 and t < max_value // point - 1:
|
||||
t += 1
|
||||
return t * point
|
||||
|
||||
|
||||
def apply_mask_strategy(z, refs_x, mask_strategys, loop_i, align=None):
|
||||
masks = []
|
||||
no_mask = True
|
||||
for i, mask_strategy in enumerate(mask_strategys):
|
||||
no_mask = False
|
||||
mask = torch.ones(z.shape[2], dtype=torch.float, device=z.device)
|
||||
mask_strategy = parse_mask_strategy(mask_strategy)
|
||||
for mst in mask_strategy:
|
||||
loop_id, m_id, m_ref_start, m_target_start, m_length, edit_ratio = mst
|
||||
if loop_id != loop_i:
|
||||
continue
|
||||
ref = refs_x[i][m_id]
|
||||
|
||||
if m_ref_start < 0:
|
||||
# ref: [C, T, H, W]
|
||||
m_ref_start = ref.shape[1] + m_ref_start
|
||||
if m_target_start < 0:
|
||||
# z: [B, C, T, H, W]
|
||||
m_target_start = z.shape[2] + m_target_start
|
||||
if align is not None:
|
||||
m_ref_start = find_nearest_point(m_ref_start, align, ref.shape[1])
|
||||
m_target_start = find_nearest_point(m_target_start, align, z.shape[2])
|
||||
m_length = min(m_length, z.shape[2] - m_target_start, ref.shape[1] - m_ref_start)
|
||||
z[i, :, m_target_start : m_target_start + m_length] = ref[:, m_ref_start : m_ref_start + m_length]
|
||||
mask[m_target_start : m_target_start + m_length] = edit_ratio
|
||||
masks.append(mask)
|
||||
if no_mask:
|
||||
return None
|
||||
masks = torch.stack(masks)
|
||||
return masks
|
||||
|
||||
|
||||
def append_generated(vae, generated_video, refs_x, mask_strategy, loop_i, condition_frame_length, condition_frame_edit):
|
||||
ref_x = vae.encode(generated_video)
|
||||
for j, refs in enumerate(refs_x):
|
||||
if refs is None:
|
||||
refs_x[j] = [ref_x[j]]
|
||||
else:
|
||||
refs.append(ref_x[j])
|
||||
if mask_strategy[j] is None or mask_strategy[j] == "":
|
||||
mask_strategy[j] = ""
|
||||
else:
|
||||
mask_strategy[j] += ";"
|
||||
mask_strategy[
|
||||
j
|
||||
] += f"{loop_i},{len(refs)-1},-{condition_frame_length},0,{condition_frame_length},{condition_frame_edit}"
|
||||
return refs_x, mask_strategy
|
||||
|
||||
|
||||
def dframe_to_frame(num):
|
||||
assert num % 5 == 0, f"Invalid num: {num}"
|
||||
return num // 5 * 17
|
||||
|
||||
|
||||
OPENAI_CLIENT = None
|
||||
REFINE_PROMPTS = None
|
||||
REFINE_PROMPTS_PATH = "assets/texts/t2v_pllava.txt"
|
||||
REFINE_PROMPTS_TEMPLATE = """
|
||||
You need to refine user's input prompt. The user's input prompt is used for video generation task. You need to refine the user's prompt to make it more suitable for the task. Here are some examples of refined prompts:
|
||||
{}
|
||||
|
||||
The refined prompt should pay attention to all objects in the video. The description should be useful for AI to re-generate the video. The description should be no more than six sentences. The refined prompt should be in English.
|
||||
"""
|
||||
RANDOM_PROMPTS = None
|
||||
RANDOM_PROMPTS_TEMPLATE = """
|
||||
You need to generate one input prompt for video generation task. The prompt should be suitable for the task. Here are some examples of refined prompts:
|
||||
{}
|
||||
|
||||
The prompt should pay attention to all objects in the video. The description should be useful for AI to re-generate the video. The description should be no more than six sentences. The prompt should be in English.
|
||||
"""
|
||||
|
||||
|
||||
def get_openai_response(sys_prompt, usr_prompt, model="gpt-4o"):
|
||||
global OPENAI_CLIENT
|
||||
if OPENAI_CLIENT is None:
|
||||
from openai import OpenAI
|
||||
|
||||
OPENAI_CLIENT = OpenAI(api_key=os.environ.get("OPENAI_API_KEY"))
|
||||
|
||||
completion = OPENAI_CLIENT.chat.completions.create(
|
||||
model=model,
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": sys_prompt,
|
||||
}, # <-- This is the system message that provides context to the model
|
||||
{
|
||||
"role": "user",
|
||||
"content": usr_prompt,
|
||||
}, # <-- This is the user message for which the model will generate a response
|
||||
],
|
||||
)
|
||||
|
||||
return completion.choices[0].message.content
|
||||
|
||||
|
||||
def get_random_prompt_by_openai():
|
||||
global RANDOM_PROMPTS
|
||||
if RANDOM_PROMPTS is None:
|
||||
examples = load_prompts(REFINE_PROMPTS_PATH)
|
||||
RANDOM_PROMPTS = RANDOM_PROMPTS_TEMPLATE.format("\n".join(examples))
|
||||
|
||||
response = get_openai_response(RANDOM_PROMPTS, "Generate one example.")
|
||||
return response
|
||||
|
||||
|
||||
def refine_prompt_by_openai(prompt):
|
||||
global REFINE_PROMPTS
|
||||
if REFINE_PROMPTS is None:
|
||||
examples = load_prompts(REFINE_PROMPTS_PATH)
|
||||
REFINE_PROMPTS = REFINE_PROMPTS_TEMPLATE.format("\n".join(examples))
|
||||
|
||||
response = get_openai_response(REFINE_PROMPTS, prompt)
|
||||
return response
|
||||
|
||||
|
||||
def has_openai_key():
|
||||
return "OPENAI_API_KEY" in os.environ
|
||||
|
||||
|
||||
def refine_prompts_by_openai(prompts):
|
||||
new_prompts = []
|
||||
for prompt in prompts:
|
||||
try:
|
||||
if prompt.strip() == "":
|
||||
new_prompt = get_random_prompt_by_openai()
|
||||
print(f"[Info] Empty prompt detected, generate random prompt: {new_prompt}")
|
||||
else:
|
||||
new_prompt = refine_prompt_by_openai(prompt)
|
||||
print(f"[Info] Refine prompt: {prompt} -> {new_prompt}")
|
||||
new_prompts.append(new_prompt)
|
||||
except Exception as e:
|
||||
print(f"[Warning] Failed to refine prompt: {prompt} due to {e}")
|
||||
new_prompts.append(prompt)
|
||||
return new_prompts
|
||||
|
||||
|
||||
def add_watermark(
|
||||
input_video_path, watermark_image_path="./assets/images/watermark/watermark.png", output_video_path=None
|
||||
):
|
||||
# execute this command in terminal with subprocess
|
||||
# return if the process is successful
|
||||
if output_video_path is None:
|
||||
output_video_path = input_video_path.replace(".mp4", "_watermark.mp4")
|
||||
cmd = f'ffmpeg -y -i {input_video_path} -i {watermark_image_path} -filter_complex "[1][0]scale2ref=oh*mdar:ih*0.1[logo][video];[video][logo]overlay" {output_video_path}'
|
||||
exit_code = os.system(cmd)
|
||||
is_success = exit_code == 0
|
||||
return is_success
|
||||
Executable
+448
@@ -0,0 +1,448 @@
|
||||
# Adapted from OpenSora
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# OpenSora: https://github.com/hpcaitech/Open-Sora
|
||||
# --------------------------------------------------------
|
||||
|
||||
import functools
|
||||
import math
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.checkpoint
|
||||
import xformers.ops
|
||||
from einops import rearrange
|
||||
from timm.models.vision_transformer import Mlp
|
||||
|
||||
approx_gelu = lambda: nn.GELU(approximate="tanh")
|
||||
|
||||
|
||||
class LlamaRMSNorm(nn.Module):
|
||||
def __init__(self, hidden_size, eps=1e-6):
|
||||
"""
|
||||
LlamaRMSNorm is equivalent to T5LayerNorm
|
||||
"""
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(hidden_size))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
def forward(self, hidden_states):
|
||||
input_dtype = hidden_states.dtype
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
variance = hidden_states.pow(2).mean(-1, keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
|
||||
return self.weight * hidden_states.to(input_dtype)
|
||||
|
||||
|
||||
def get_layernorm(hidden_size: torch.Tensor, eps: float, affine: bool):
|
||||
return nn.LayerNorm(hidden_size, eps, elementwise_affine=affine)
|
||||
|
||||
|
||||
def t2i_modulate(x, shift, scale):
|
||||
return x * (1 + scale) + shift
|
||||
|
||||
|
||||
# ===============================================
|
||||
# General-purpose Layers
|
||||
# ===============================================
|
||||
|
||||
|
||||
class PatchEmbed3D(nn.Module):
|
||||
"""Video to Patch Embedding.
|
||||
|
||||
Args:
|
||||
patch_size (int): Patch token size. Default: (2,4,4).
|
||||
in_chans (int): Number of input video channels. Default: 3.
|
||||
embed_dim (int): Number of linear projection output channels. Default: 96.
|
||||
norm_layer (nn.Module, optional): Normalization layer. Default: None
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=(2, 4, 4),
|
||||
in_chans=3,
|
||||
embed_dim=96,
|
||||
norm_layer=None,
|
||||
flatten=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.flatten = flatten
|
||||
|
||||
self.in_chans = in_chans
|
||||
self.embed_dim = embed_dim
|
||||
|
||||
self.proj = nn.Conv3d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
|
||||
if norm_layer is not None:
|
||||
self.norm = norm_layer(embed_dim)
|
||||
else:
|
||||
self.norm = None
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward function."""
|
||||
# padding
|
||||
_, _, D, H, W = x.size()
|
||||
if W % self.patch_size[2] != 0:
|
||||
x = F.pad(x, (0, self.patch_size[2] - W % self.patch_size[2]))
|
||||
if H % self.patch_size[1] != 0:
|
||||
x = F.pad(x, (0, 0, 0, self.patch_size[1] - H % self.patch_size[1]))
|
||||
if D % self.patch_size[0] != 0:
|
||||
x = F.pad(x, (0, 0, 0, 0, 0, self.patch_size[0] - D % self.patch_size[0]))
|
||||
|
||||
x = self.proj(x) # (B C T H W)
|
||||
if self.norm is not None:
|
||||
D, Wh, Ww = x.size(2), x.size(3), x.size(4)
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
x = self.norm(x)
|
||||
x = x.transpose(1, 2).view(-1, self.embed_dim, D, Wh, Ww)
|
||||
if self.flatten:
|
||||
x = x.flatten(2).transpose(1, 2) # BCTHW -> BNC
|
||||
return x
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_heads: int = 8,
|
||||
qkv_bias: bool = False,
|
||||
qk_norm: bool = False,
|
||||
attn_drop: float = 0.0,
|
||||
proj_drop: float = 0.0,
|
||||
norm_layer: nn.Module = LlamaRMSNorm,
|
||||
enable_flash_attn: bool = False,
|
||||
rope=None,
|
||||
qk_norm_legacy: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
assert dim % num_heads == 0, "dim should be divisible by num_heads"
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.scale = self.head_dim**-0.5
|
||||
self.enable_flash_attn = enable_flash_attn
|
||||
|
||||
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
self.q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
|
||||
self.k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
|
||||
self.qk_norm_legacy = qk_norm_legacy
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
|
||||
self.rope = False
|
||||
if rope is not None:
|
||||
self.rope = True
|
||||
self.rotary_emb = rope
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
B, N, C = x.shape
|
||||
# flash attn is not memory efficient for small sequences, this is empirical
|
||||
enable_flash_attn = self.enable_flash_attn and (N > B)
|
||||
qkv = self.qkv(x)
|
||||
qkv_shape = (B, N, 3, self.num_heads, self.head_dim)
|
||||
|
||||
qkv = qkv.view(qkv_shape).permute(2, 0, 3, 1, 4)
|
||||
q, k, v = qkv.unbind(0)
|
||||
if self.qk_norm_legacy:
|
||||
# WARNING: this may be a bug
|
||||
if self.rope:
|
||||
q = self.rotary_emb(q)
|
||||
k = self.rotary_emb(k)
|
||||
q, k = self.q_norm(q), self.k_norm(k)
|
||||
else:
|
||||
q, k = self.q_norm(q), self.k_norm(k)
|
||||
if self.rope:
|
||||
q = self.rotary_emb(q)
|
||||
k = self.rotary_emb(k)
|
||||
|
||||
if enable_flash_attn:
|
||||
from flash_attn import flash_attn_func
|
||||
|
||||
# (B, #heads, N, #dim) -> (B, N, #heads, #dim)
|
||||
q = q.permute(0, 2, 1, 3)
|
||||
k = k.permute(0, 2, 1, 3)
|
||||
v = v.permute(0, 2, 1, 3)
|
||||
x = flash_attn_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=self.attn_drop.p if self.training else 0.0,
|
||||
softmax_scale=self.scale,
|
||||
)
|
||||
else:
|
||||
dtype = q.dtype
|
||||
q = q * self.scale
|
||||
attn = q @ k.transpose(-2, -1) # translate attn to float32
|
||||
attn = attn.to(torch.float32)
|
||||
attn = attn.softmax(dim=-1)
|
||||
attn = attn.to(dtype) # cast back attn to original dtype
|
||||
attn = self.attn_drop(attn)
|
||||
x = attn @ v
|
||||
|
||||
x_output_shape = (B, N, C)
|
||||
if not enable_flash_attn:
|
||||
x = x.transpose(1, 2)
|
||||
x = x.reshape(x_output_shape)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class MultiHeadCrossAttention(nn.Module):
|
||||
def __init__(self, d_model, num_heads, attn_drop=0.0, proj_drop=0.0):
|
||||
super(MultiHeadCrossAttention, self).__init__()
|
||||
assert d_model % num_heads == 0, "d_model must be divisible by num_heads"
|
||||
|
||||
self.d_model = d_model
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = d_model // num_heads
|
||||
|
||||
self.q_linear = nn.Linear(d_model, d_model)
|
||||
self.kv_linear = nn.Linear(d_model, d_model * 2)
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(d_model, d_model)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
|
||||
def forward(self, x, cond, mask=None):
|
||||
# query/value: img tokens; key: condition; mask: if padding tokens
|
||||
B, N, C = x.shape
|
||||
|
||||
q = self.q_linear(x).view(1, -1, self.num_heads, self.head_dim)
|
||||
kv = self.kv_linear(cond).view(1, -1, 2, self.num_heads, self.head_dim)
|
||||
k, v = kv.unbind(2)
|
||||
|
||||
attn_bias = None
|
||||
if mask is not None:
|
||||
attn_bias = xformers.ops.fmha.BlockDiagonalMask.from_seqlens([N] * B, mask)
|
||||
x = xformers.ops.memory_efficient_attention(q, k, v, p=self.attn_drop.p, attn_bias=attn_bias)
|
||||
|
||||
x = x.view(B, -1, C)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class T2IFinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, num_patch, out_channels, d_t=None, d_s=None):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size, num_patch * out_channels, bias=True)
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(2, hidden_size) / hidden_size**0.5)
|
||||
self.out_channels = out_channels
|
||||
self.d_t = d_t
|
||||
self.d_s = d_s
|
||||
|
||||
def t_mask_select(self, x_mask, x, masked_x, T, S):
|
||||
# x: [B, (T, S), C]
|
||||
# mased_x: [B, (T, S), C]
|
||||
# x_mask: [B, T]
|
||||
x = rearrange(x, "B (T S) C -> B T S C", T=T, S=S)
|
||||
masked_x = rearrange(masked_x, "B (T S) C -> B T S C", T=T, S=S)
|
||||
x = torch.where(x_mask[:, :, None, None], x, masked_x)
|
||||
x = rearrange(x, "B T S C -> B (T S) C")
|
||||
return x
|
||||
|
||||
def forward(self, x, t, x_mask=None, t0=None, T=None, S=None):
|
||||
if T is None:
|
||||
T = self.d_t
|
||||
if S is None:
|
||||
S = self.d_s
|
||||
shift, scale = (self.scale_shift_table[None] + t[:, None]).chunk(2, dim=1)
|
||||
x = t2i_modulate(self.norm_final(x), shift, scale)
|
||||
if x_mask is not None:
|
||||
shift_zero, scale_zero = (self.scale_shift_table[None] + t0[:, None]).chunk(2, dim=1)
|
||||
x_zero = t2i_modulate(self.norm_final(x), shift_zero, scale_zero)
|
||||
x = self.t_mask_select(x_mask, x, x_zero, T, S)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
# ===============================================
|
||||
# Embedding Layers for Timesteps and Class Labels
|
||||
# ===============================================
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__()
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
||||
)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
|
||||
@staticmethod
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param t: 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, D) Tensor of positional embeddings.
|
||||
"""
|
||||
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
||||
half = dim // 2
|
||||
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half)
|
||||
freqs = freqs.to(device=t.device)
|
||||
args = t[:, 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)
|
||||
return embedding
|
||||
|
||||
def forward(self, t, dtype):
|
||||
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
||||
if t_freq.dtype != dtype:
|
||||
t_freq = t_freq.to(dtype)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
|
||||
|
||||
class SizeEmbedder(TimestepEmbedder):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__(hidden_size=hidden_size, frequency_embedding_size=frequency_embedding_size)
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
||||
)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
self.outdim = hidden_size
|
||||
|
||||
def forward(self, s, bs):
|
||||
if s.ndim == 1:
|
||||
s = s[:, None]
|
||||
assert s.ndim == 2
|
||||
if s.shape[0] != bs:
|
||||
s = s.repeat(bs // s.shape[0], 1)
|
||||
assert s.shape[0] == bs
|
||||
b, dims = s.shape[0], s.shape[1]
|
||||
s = rearrange(s, "b d -> (b d)")
|
||||
s_freq = self.timestep_embedding(s, self.frequency_embedding_size).to(self.dtype)
|
||||
s_emb = self.mlp(s_freq)
|
||||
s_emb = rearrange(s_emb, "(b d) d2 -> b (d d2)", b=b, d=dims, d2=self.outdim)
|
||||
return s_emb
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
|
||||
class CaptionEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
hidden_size,
|
||||
uncond_prob,
|
||||
act_layer=nn.GELU(approximate="tanh"),
|
||||
token_num=120,
|
||||
):
|
||||
super().__init__()
|
||||
self.y_proj = Mlp(
|
||||
in_features=in_channels,
|
||||
hidden_features=hidden_size,
|
||||
out_features=hidden_size,
|
||||
act_layer=act_layer,
|
||||
drop=0,
|
||||
)
|
||||
self.register_buffer(
|
||||
"y_embedding",
|
||||
torch.randn(token_num, in_channels) / in_channels**0.5,
|
||||
)
|
||||
self.uncond_prob = uncond_prob
|
||||
|
||||
def token_drop(self, caption, force_drop_ids=None):
|
||||
"""
|
||||
Drops labels to enable classifier-free guidance.
|
||||
"""
|
||||
if force_drop_ids is None:
|
||||
drop_ids = torch.rand(caption.shape[0]).cuda() < self.uncond_prob
|
||||
else:
|
||||
drop_ids = force_drop_ids == 1
|
||||
caption = torch.where(drop_ids[:, None, None, None], self.y_embedding, caption)
|
||||
return caption
|
||||
|
||||
def forward(self, caption, train, force_drop_ids=None):
|
||||
if train:
|
||||
assert caption.shape[2:] == self.y_embedding.shape
|
||||
use_dropout = self.uncond_prob > 0
|
||||
if (train and use_dropout) or (force_drop_ids is not None):
|
||||
caption = self.token_drop(caption, force_drop_ids)
|
||||
caption = self.y_proj(caption)
|
||||
return caption
|
||||
|
||||
|
||||
class PositionEmbedding2D(nn.Module):
|
||||
def __init__(self, dim: int) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
assert dim % 4 == 0, "dim must be divisible by 4"
|
||||
half_dim = dim // 2
|
||||
inv_freq = 1.0 / (10000 ** (torch.arange(0, half_dim, 2).float() / half_dim))
|
||||
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
||||
|
||||
def _get_sin_cos_emb(self, t: torch.Tensor):
|
||||
out = torch.einsum("i,d->id", t, self.inv_freq)
|
||||
emb_cos = torch.cos(out)
|
||||
emb_sin = torch.sin(out)
|
||||
return torch.cat((emb_sin, emb_cos), dim=-1)
|
||||
|
||||
@functools.lru_cache(maxsize=512)
|
||||
def _get_cached_emb(
|
||||
self,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
h: int,
|
||||
w: int,
|
||||
scale: float = 1.0,
|
||||
base_size: Optional[int] = None,
|
||||
):
|
||||
grid_h = torch.arange(h, device=device) / scale
|
||||
grid_w = torch.arange(w, device=device) / scale
|
||||
if base_size is not None:
|
||||
grid_h *= base_size / h
|
||||
grid_w *= base_size / w
|
||||
grid_h, grid_w = torch.meshgrid(
|
||||
grid_w,
|
||||
grid_h,
|
||||
indexing="ij",
|
||||
) # here w goes first
|
||||
grid_h = grid_h.t().reshape(-1)
|
||||
grid_w = grid_w.t().reshape(-1)
|
||||
emb_h = self._get_sin_cos_emb(grid_h)
|
||||
emb_w = self._get_sin_cos_emb(grid_w)
|
||||
return torch.concat([emb_h, emb_w], dim=-1).unsqueeze(0).to(dtype)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
h: int,
|
||||
w: int,
|
||||
scale: Optional[float] = 1.0,
|
||||
base_size: Optional[int] = None,
|
||||
) -> torch.Tensor:
|
||||
return self._get_cached_emb(x.device, x.dtype, h, w, scale, base_size)
|
||||
Executable
+267
@@ -0,0 +1,267 @@
|
||||
# Adapted from OpenSora
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# OpenSora: https://github.com/hpcaitech/Open-Sora
|
||||
# --------------------------------------------------------
|
||||
|
||||
import torch
|
||||
#import torch.distributed as dist
|
||||
from einops import rearrange
|
||||
#from torch.distributions import LogisticNormal
|
||||
from tqdm import tqdm
|
||||
|
||||
from opendit.core.pab_mgr import get_diffusion_skip, get_diffusion_skip_timestep, skip_diffusion_timestep
|
||||
from opendit.diffusion.gaussian_diffusion import _extract_into_tensor
|
||||
|
||||
from comfy.utils import ProgressBar
|
||||
def mean_flat(tensor: torch.Tensor, mask=None):
|
||||
"""
|
||||
Take the mean over all non-batch dimensions.
|
||||
"""
|
||||
if mask is None:
|
||||
return tensor.mean(dim=list(range(1, len(tensor.shape))))
|
||||
else:
|
||||
assert tensor.dim() == 5
|
||||
assert tensor.shape[2] == mask.shape[1]
|
||||
tensor = rearrange(tensor, "b c t h w -> b t (c h w)")
|
||||
denom = mask.sum(dim=1) * tensor.shape[-1]
|
||||
loss = (tensor * mask.unsqueeze(2)).sum(dim=1).sum(dim=1) / denom
|
||||
return loss
|
||||
|
||||
|
||||
def timestep_transform(
|
||||
t,
|
||||
model_kwargs,
|
||||
base_resolution=512 * 512,
|
||||
base_num_frames=1,
|
||||
scale=1.0,
|
||||
num_timesteps=1,
|
||||
):
|
||||
t = t / num_timesteps
|
||||
resolution = model_kwargs["height"] * model_kwargs["width"]
|
||||
ratio_space = (resolution / base_resolution).sqrt()
|
||||
# NOTE: currently, we do not take fps into account
|
||||
# NOTE: temporal_reduction is hardcoded, this should be equal to the temporal reduction factor of the vae
|
||||
if model_kwargs["num_frames"][0] == 1:
|
||||
num_frames = torch.ones_like(model_kwargs["num_frames"])
|
||||
else:
|
||||
num_frames = model_kwargs["num_frames"] // 17 * 5
|
||||
ratio_time = (num_frames / base_num_frames).sqrt()
|
||||
|
||||
ratio = ratio_space * ratio_time * scale
|
||||
new_t = ratio * t / (1 + (ratio - 1) * t)
|
||||
|
||||
new_t = new_t * num_timesteps
|
||||
return new_t
|
||||
|
||||
|
||||
class RFlowScheduler:
|
||||
def __init__(
|
||||
self,
|
||||
num_timesteps=1000,
|
||||
num_sampling_steps=10,
|
||||
use_discrete_timesteps=False,
|
||||
sample_method="uniform",
|
||||
loc=0.0,
|
||||
scale=1.0,
|
||||
use_timestep_transform=False,
|
||||
transform_scale=1.0,
|
||||
):
|
||||
self.num_timesteps = num_timesteps
|
||||
self.num_sampling_steps = num_sampling_steps
|
||||
self.use_discrete_timesteps = use_discrete_timesteps
|
||||
|
||||
# sample method
|
||||
assert sample_method in ["uniform", "logit-normal"]
|
||||
assert (
|
||||
sample_method == "uniform" or not use_discrete_timesteps
|
||||
), "Only uniform sampling is supported for discrete timesteps"
|
||||
self.sample_method = sample_method
|
||||
if sample_method == "logit-normal":
|
||||
self.distribution = LogisticNormal(torch.tensor([loc]), torch.tensor([scale]))
|
||||
self.sample_t = lambda x: self.distribution.sample((x.shape[0],))[:, 0].to(x.device)
|
||||
|
||||
# timestep transform
|
||||
self.use_timestep_transform = use_timestep_transform
|
||||
self.transform_scale = transform_scale
|
||||
|
||||
def training_losses(self, model, x_start, model_kwargs=None, noise=None, mask=None, weights=None, t=None):
|
||||
"""
|
||||
Compute training losses for a single timestep.
|
||||
Arguments format copied from opensora/schedulers/iddpm/gaussian_diffusion.py/training_losses
|
||||
Note: t is int tensor and should be rescaled from [0, num_timesteps-1] to [1,0]
|
||||
"""
|
||||
if t is None:
|
||||
if self.use_discrete_timesteps:
|
||||
t = torch.randint(0, self.num_timesteps, (x_start.shape[0],), device=x_start.device)
|
||||
elif self.sample_method == "uniform":
|
||||
t = torch.rand((x_start.shape[0],), device=x_start.device) * self.num_timesteps
|
||||
elif self.sample_method == "logit-normal":
|
||||
t = self.sample_t(x_start) * self.num_timesteps
|
||||
|
||||
if self.use_timestep_transform:
|
||||
t = timestep_transform(t, model_kwargs, scale=self.transform_scale, num_timesteps=self.num_timesteps)
|
||||
|
||||
if model_kwargs is None:
|
||||
model_kwargs = {}
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x_start)
|
||||
assert noise.shape == x_start.shape
|
||||
|
||||
x_t = self.add_noise(x_start, noise, t)
|
||||
if mask is not None:
|
||||
t0 = torch.zeros_like(t)
|
||||
x_t0 = self.add_noise(x_start, noise, t0)
|
||||
x_t = torch.where(mask[:, None, :, None, None], x_t, x_t0)
|
||||
|
||||
terms = {}
|
||||
model_output = model(x_t, t, **model_kwargs)
|
||||
velocity_pred = model_output.chunk(2, dim=1)[0]
|
||||
if weights is None:
|
||||
loss = mean_flat((velocity_pred - (x_start - noise)).pow(2), mask=mask)
|
||||
else:
|
||||
weight = _extract_into_tensor(weights, t, x_start.shape)
|
||||
loss = mean_flat(weight * (velocity_pred - (x_start - noise)).pow(2), mask=mask)
|
||||
terms["loss"] = loss
|
||||
|
||||
return terms
|
||||
|
||||
def add_noise(
|
||||
self,
|
||||
original_samples: torch.FloatTensor,
|
||||
noise: torch.FloatTensor,
|
||||
timesteps: torch.IntTensor,
|
||||
) -> torch.FloatTensor:
|
||||
"""
|
||||
compatible with diffusers add_noise()
|
||||
"""
|
||||
timepoints = timesteps.float() / self.num_timesteps
|
||||
timepoints = 1 - timepoints # [1,1/1000]
|
||||
|
||||
# timepoint (bsz) noise: (bsz, 4, frame, w ,h)
|
||||
# expand timepoint to noise shape
|
||||
timepoints = timepoints.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1)
|
||||
timepoints = timepoints.repeat(1, noise.shape[1], noise.shape[2], noise.shape[3], noise.shape[4])
|
||||
|
||||
return timepoints * original_samples + (1 - timepoints) * noise
|
||||
|
||||
|
||||
class RFLOW:
|
||||
def __init__(
|
||||
self,
|
||||
num_sampling_steps=10,
|
||||
num_timesteps=1000,
|
||||
cfg_scale=4.0,
|
||||
use_discrete_timesteps=False,
|
||||
use_timestep_transform=False,
|
||||
**kwargs,
|
||||
):
|
||||
self.num_sampling_steps = num_sampling_steps
|
||||
self.num_timesteps = num_timesteps
|
||||
self.cfg_scale = cfg_scale
|
||||
self.use_discrete_timesteps = use_discrete_timesteps
|
||||
self.use_timestep_transform = use_timestep_transform
|
||||
|
||||
self.scheduler = RFlowScheduler(
|
||||
num_timesteps=num_timesteps,
|
||||
num_sampling_steps=num_sampling_steps,
|
||||
use_discrete_timesteps=use_discrete_timesteps,
|
||||
use_timestep_transform=use_timestep_transform,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def sample(
|
||||
self,
|
||||
model,
|
||||
model_args,
|
||||
z,
|
||||
device,
|
||||
additional_args=None,
|
||||
mask=None,
|
||||
guidance_scale=None,
|
||||
progress=True,
|
||||
verbose=False,
|
||||
):
|
||||
# if no specific guidance scale is provided, use the default scale when initializing the scheduler
|
||||
if guidance_scale is None:
|
||||
guidance_scale = self.cfg_scale
|
||||
|
||||
#n = len(prompts)
|
||||
# text encoding
|
||||
# model_args = text_encoder.encode(prompts)
|
||||
# #y_null = text_encoder.null(n)
|
||||
# y_null = y_embedder.y_embedding[None].repeat(n, 1, 1)[:, None]
|
||||
# model_args["y"] = torch.cat([model_args["y"], y_null], 0)
|
||||
# if additional_args is not None:
|
||||
# model_args.update(additional_args)
|
||||
|
||||
# prepare timesteps
|
||||
timesteps = [(1.0 - i / self.num_sampling_steps) * self.num_timesteps for i in range(self.num_sampling_steps)]
|
||||
if self.use_discrete_timesteps:
|
||||
timesteps = [int(round(t)) for t in timesteps]
|
||||
timesteps = [torch.tensor([t] * z.shape[0], device=device) for t in timesteps]
|
||||
if self.use_timestep_transform:
|
||||
timesteps = [timestep_transform(t, additional_args, num_timesteps=self.num_timesteps) for t in timesteps]
|
||||
# print(f'timesteps: {timesteps}')
|
||||
# TODO: jump diffusion steps
|
||||
|
||||
if get_diffusion_skip() and get_diffusion_skip_timestep() is not None:
|
||||
orignal_timesteps = timesteps
|
||||
diffusion_skip_timestep = get_diffusion_skip_timestep()
|
||||
timesteps = skip_diffusion_timestep(timesteps, diffusion_skip_timestep)
|
||||
|
||||
if verbose:
|
||||
print("============================")
|
||||
print("skip diffusion steps!!!")
|
||||
print("============================")
|
||||
print(f"orignal sample timesteps: {orignal_timesteps}")
|
||||
print(f"orignal diffusion steps: {len(orignal_timesteps)}")
|
||||
print("============================")
|
||||
print(f"skip diffusion steps: {get_diffusion_skip_timestep()}")
|
||||
print(f"sample timesteps: {timesteps}")
|
||||
print(f"num_inference_steps: {len(timesteps)}")
|
||||
print("============================")
|
||||
|
||||
if mask is not None:
|
||||
noise_added = torch.zeros_like(mask, dtype=torch.bool)
|
||||
noise_added = noise_added | (mask == 1)
|
||||
|
||||
comfy_pbar=ProgressBar(len(timesteps))
|
||||
progress_wrap = tqdm if progress else (lambda x: x)
|
||||
for i, t in progress_wrap(list(enumerate(timesteps))):
|
||||
# mask for adding noise
|
||||
if mask is not None:
|
||||
mask_t = mask * self.num_timesteps
|
||||
x0 = z.clone()
|
||||
x_noise = self.scheduler.add_noise(x0, torch.randn_like(x0), t)
|
||||
|
||||
mask_t_upper = mask_t >= t.unsqueeze(1)
|
||||
model_args["x_mask"] = mask_t_upper.repeat(2, 1)
|
||||
mask_add_noise = mask_t_upper & ~noise_added
|
||||
|
||||
z = torch.where(mask_add_noise[:, None, :, None, None], x_noise, x0)
|
||||
noise_added = mask_t_upper
|
||||
|
||||
# classifier-free guidance
|
||||
z_in = torch.cat([z, z], 0)
|
||||
t = torch.cat([t, t], 0)
|
||||
pred = model(z_in, t, **model_args).chunk(2, dim=1)[0]
|
||||
pred_cond, pred_uncond = pred.chunk(2, dim=0)
|
||||
v_pred = pred_uncond + guidance_scale * (pred_cond - pred_uncond)
|
||||
|
||||
# update z
|
||||
dt = timesteps[i] - timesteps[i + 1] if i < len(timesteps) - 1 else timesteps[i]
|
||||
dt = dt / self.num_timesteps
|
||||
z = z + v_pred * dt[:, None, None, None, None]
|
||||
|
||||
if mask is not None:
|
||||
z = torch.where(mask_t_upper[:, None, :, None, None], z, x0)
|
||||
comfy_pbar.update(1)
|
||||
|
||||
return z
|
||||
|
||||
def training_losses(self, model, x_start, model_kwargs=None, noise=None, mask=None, weights=None, t=None):
|
||||
return self.scheduler.training_losses(model, x_start, model_kwargs, noise, mask, weights, t)
|
||||
Executable
+702
@@ -0,0 +1,702 @@
|
||||
# Adapted from OpenSora
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# OpenSora: https://github.com/hpcaitech/Open-Sora
|
||||
# --------------------------------------------------------
|
||||
|
||||
import functools
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from timm.models.layers import DropPath
|
||||
from timm.models.vision_transformer import Mlp
|
||||
from transformers import PretrainedConfig, PreTrainedModel
|
||||
|
||||
from opendit.core.comm import all_to_all_comm, gather_sequence, split_sequence
|
||||
from opendit.core.parallel_mgr import (
|
||||
get_sequence_parallel_group,
|
||||
get_sequence_parallel_size,
|
||||
is_sequence_parallelism_enable,
|
||||
)
|
||||
from opendit.models.opensora.ckpt_io import load_checkpoint
|
||||
from opendit.models.opensora.embed import CaptionEmbedder, PatchEmbed3D, TimestepEmbedder, get_2d_sincos_pos_embed
|
||||
from opendit.models.opensora.stdit import approx_gelu, t2i_modulate
|
||||
from opendit.modules.attn import Attention, MultiHeadCrossAttention
|
||||
from opendit.modules.layers import get_layernorm
|
||||
|
||||
|
||||
class PositionEmbedding2D(nn.Module):
|
||||
def __init__(self, dim: int) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
assert dim % 4 == 0, "dim must be divisible by 4"
|
||||
half_dim = dim // 2
|
||||
inv_freq = 1.0 / (10000 ** (torch.arange(0, half_dim, 2).float() / half_dim))
|
||||
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
||||
|
||||
def _get_sin_cos_emb(self, t: torch.Tensor):
|
||||
out = torch.einsum("i,d->id", t, self.inv_freq)
|
||||
emb_cos = torch.cos(out)
|
||||
emb_sin = torch.sin(out)
|
||||
return torch.cat((emb_sin, emb_cos), dim=-1)
|
||||
|
||||
@functools.lru_cache(maxsize=512)
|
||||
def _get_cached_emb(
|
||||
self,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
h: int,
|
||||
w: int,
|
||||
scale: float = 1.0,
|
||||
base_size: Optional[int] = None,
|
||||
):
|
||||
grid_h = torch.arange(h, device=device) / scale
|
||||
grid_w = torch.arange(w, device=device) / scale
|
||||
if base_size is not None:
|
||||
grid_h *= base_size / h
|
||||
grid_w *= base_size / w
|
||||
grid_h, grid_w = torch.meshgrid(
|
||||
grid_w,
|
||||
grid_h,
|
||||
indexing="ij",
|
||||
) # here w goes first
|
||||
grid_h = grid_h.t().reshape(-1)
|
||||
grid_w = grid_w.t().reshape(-1)
|
||||
emb_h = self._get_sin_cos_emb(grid_h)
|
||||
emb_w = self._get_sin_cos_emb(grid_w)
|
||||
return torch.concat([emb_h, emb_w], dim=-1).unsqueeze(0).to(dtype)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
h: int,
|
||||
w: int,
|
||||
scale: Optional[float] = 1.0,
|
||||
base_size: Optional[int] = None,
|
||||
) -> torch.Tensor:
|
||||
return self._get_cached_emb(x.device, x.dtype, h, w, scale, base_size)
|
||||
|
||||
|
||||
class SizeEmbedder(TimestepEmbedder):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__(hidden_size=hidden_size, frequency_embedding_size=frequency_embedding_size)
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
||||
)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
self.outdim = hidden_size
|
||||
|
||||
def forward(self, s, bs):
|
||||
if s.ndim == 1:
|
||||
s = s[:, None]
|
||||
assert s.ndim == 2
|
||||
if s.shape[0] != bs:
|
||||
s = s.repeat(bs // s.shape[0], 1)
|
||||
assert s.shape[0] == bs
|
||||
b, dims = s.shape[0], s.shape[1]
|
||||
s = rearrange(s, "b d -> (b d)")
|
||||
s_freq = self.timestep_embedding(s, self.frequency_embedding_size).to(self.dtype)
|
||||
s_emb = self.mlp(s_freq)
|
||||
s_emb = rearrange(s_emb, "(b d) d2 -> b (d d2)", b=b, d=dims, d2=self.outdim)
|
||||
return s_emb
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
|
||||
class T2IFinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, num_patch, out_channels, d_t=None, d_s=None):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size, num_patch * out_channels, bias=True)
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(2, hidden_size) / hidden_size**0.5)
|
||||
self.out_channels = out_channels
|
||||
self.d_t = d_t
|
||||
self.d_s = d_s
|
||||
|
||||
def t_mask_select(self, x_mask, x, masked_x, T, S):
|
||||
# x: [B, (T, S), C]
|
||||
# mased_x: [B, (T, S), C]
|
||||
# x_mask: [B, T]
|
||||
x = rearrange(x, "B (T S) C -> B T S C", T=T, S=S)
|
||||
masked_x = rearrange(masked_x, "B (T S) C -> B T S C", T=T, S=S)
|
||||
x = torch.where(x_mask[:, :, None, None], x, masked_x)
|
||||
x = rearrange(x, "B T S C -> B (T S) C")
|
||||
return x
|
||||
|
||||
def forward(self, x, t, x_mask=None, t0=None, T=None, S=None):
|
||||
if T is None:
|
||||
T = self.d_t
|
||||
if S is None:
|
||||
S = self.d_s
|
||||
shift, scale = (self.scale_shift_table[None] + t[:, None]).chunk(2, dim=1)
|
||||
x = t2i_modulate(self.norm_final(x), shift, scale)
|
||||
if x_mask is not None:
|
||||
shift_zero, scale_zero = (self.scale_shift_table[None] + t0[:, None]).chunk(2, dim=1)
|
||||
x_zero = t2i_modulate(self.norm_final(x), shift_zero, scale_zero)
|
||||
x = self.t_mask_select(x_mask, x, x_zero, T, S)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class STDiT2Block(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
num_heads,
|
||||
mlp_ratio=4.0,
|
||||
drop_path=0.0,
|
||||
enable_flash_attn=False,
|
||||
enable_layernorm_kernel=False,
|
||||
enable_sequence_parallelism=False,
|
||||
rope=None,
|
||||
qk_norm=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.enable_flash_attn = enable_flash_attn
|
||||
self._enable_sequence_parallelism = enable_sequence_parallelism
|
||||
|
||||
# spatial branch
|
||||
self.norm1 = get_layernorm(hidden_size, eps=1e-6, affine=False, use_kernel=enable_layernorm_kernel)
|
||||
self.attn = Attention(
|
||||
hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
enable_flashattn=enable_flash_attn,
|
||||
qk_norm=qk_norm,
|
||||
)
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(6, hidden_size) / hidden_size**0.5)
|
||||
|
||||
# cross attn
|
||||
self.cross_attn = MultiHeadCrossAttention(hidden_size, num_heads, enable_flashattn=enable_flash_attn)
|
||||
|
||||
# mlp branch
|
||||
self.norm2 = get_layernorm(hidden_size, eps=1e-6, affine=False, use_kernel=enable_layernorm_kernel)
|
||||
self.mlp = Mlp(
|
||||
in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0
|
||||
)
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
|
||||
|
||||
# temporal branch
|
||||
self.norm_temp = get_layernorm(hidden_size, eps=1e-6, affine=False, use_kernel=enable_layernorm_kernel) # new
|
||||
self.attn_temp = Attention(
|
||||
hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
enable_flashattn=self.enable_flash_attn,
|
||||
rope=rope,
|
||||
qk_norm=qk_norm,
|
||||
)
|
||||
self.scale_shift_table_temporal = nn.Parameter(torch.randn(3, hidden_size) / hidden_size**0.5) # new
|
||||
|
||||
def t_mask_select(self, x_mask, x, masked_x, T, S):
|
||||
# x: [B, (T, S), C]
|
||||
# mased_x: [B, (T, S), C]
|
||||
# x_mask: [B, T]
|
||||
x = rearrange(x, "B (T S) C -> B T S C", T=T, S=S)
|
||||
masked_x = rearrange(masked_x, "B (T S) C -> B T S C", T=T, S=S)
|
||||
x = torch.where(x_mask[:, :, None, None], x, masked_x)
|
||||
x = rearrange(x, "B T S C -> B (T S) C")
|
||||
return x
|
||||
|
||||
def forward(self, x, y, t, t_tmp, mask=None, x_mask=None, t0=None, t0_tmp=None, T=None, S=None):
|
||||
B, N, C = x.shape
|
||||
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||
self.scale_shift_table[None] + t.reshape(B, 6, -1)
|
||||
).chunk(6, dim=1)
|
||||
shift_tmp, scale_tmp, gate_tmp = (self.scale_shift_table_temporal[None] + t_tmp.reshape(B, 3, -1)).chunk(
|
||||
3, dim=1
|
||||
)
|
||||
if x_mask is not None:
|
||||
shift_msa_zero, scale_msa_zero, gate_msa_zero, shift_mlp_zero, scale_mlp_zero, gate_mlp_zero = (
|
||||
self.scale_shift_table[None] + t0.reshape(B, 6, -1)
|
||||
).chunk(6, dim=1)
|
||||
shift_tmp_zero, scale_tmp_zero, gate_tmp_zero = (
|
||||
self.scale_shift_table_temporal[None] + t0_tmp.reshape(B, 3, -1)
|
||||
).chunk(3, dim=1)
|
||||
|
||||
# modulate
|
||||
x_m = t2i_modulate(self.norm1(x), shift_msa, scale_msa)
|
||||
if x_mask is not None:
|
||||
x_m_zero = t2i_modulate(self.norm1(x), shift_msa_zero, scale_msa_zero)
|
||||
x_m = self.t_mask_select(x_mask, x_m, x_m_zero, T, S)
|
||||
|
||||
# spatial branch
|
||||
x_s = rearrange(x_m, "B (T S) C -> (B T) S C", T=T, S=S)
|
||||
x_s = self.attn(x_s)
|
||||
x_s = rearrange(x_s, "(B T) S C -> B (T S) C", T=T, S=S)
|
||||
if x_mask is not None:
|
||||
x_s_zero = gate_msa_zero * x_s
|
||||
x_s = gate_msa * x_s
|
||||
x_s = self.t_mask_select(x_mask, x_s, x_s_zero, T, S)
|
||||
else:
|
||||
x_s = gate_msa * x_s
|
||||
x = x + self.drop_path(x_s)
|
||||
|
||||
# modulate
|
||||
x_m = t2i_modulate(self.norm_temp(x), shift_tmp, scale_tmp)
|
||||
if x_mask is not None:
|
||||
x_m_zero = t2i_modulate(self.norm_temp(x), shift_tmp_zero, scale_tmp_zero)
|
||||
x_m = self.t_mask_select(x_mask, x_m, x_m_zero, T, S)
|
||||
|
||||
# temporal branch
|
||||
if is_sequence_parallelism_enable():
|
||||
x_m, S, T = self.dynamic_switch(x_m, S, T, temporal_to_spatial=True)
|
||||
|
||||
x_t = rearrange(x_m, "B (T S) C -> (B S) T C", T=T, S=S)
|
||||
x_t = self.attn_temp(x_t)
|
||||
x_t = rearrange(x_t, "(B S) T C -> B (T S) C", T=T, S=S)
|
||||
|
||||
if is_sequence_parallelism_enable():
|
||||
x_t, S, T = self.dynamic_switch(x_t, S, T, temporal_to_spatial=False)
|
||||
|
||||
if x_mask is not None:
|
||||
x_t_zero = gate_tmp_zero * x_t
|
||||
x_t = gate_tmp * x_t
|
||||
x_t = self.t_mask_select(x_mask, x_t, x_t_zero, T, S)
|
||||
else:
|
||||
x_t = gate_tmp * x_t
|
||||
x = x + self.drop_path(x_t)
|
||||
|
||||
# cross attn
|
||||
x = x + self.cross_attn(x, y, mask)
|
||||
|
||||
# modulate
|
||||
x_m = t2i_modulate(self.norm2(x), shift_mlp, scale_mlp)
|
||||
if x_mask is not None:
|
||||
x_m_zero = t2i_modulate(self.norm2(x), shift_mlp_zero, scale_mlp_zero)
|
||||
x_m = self.t_mask_select(x_mask, x_m, x_m_zero, T, S)
|
||||
|
||||
# mlp
|
||||
x_mlp = self.mlp(x_m)
|
||||
if x_mask is not None:
|
||||
x_mlp_zero = gate_mlp_zero * x_mlp
|
||||
x_mlp = gate_mlp * x_mlp
|
||||
x_mlp = self.t_mask_select(x_mask, x_mlp, x_mlp_zero, T, S)
|
||||
else:
|
||||
x_mlp = gate_mlp * x_mlp
|
||||
x = x + self.drop_path(x_mlp)
|
||||
|
||||
return x
|
||||
|
||||
def dynamic_switch(self, x, s, t, temporal_to_spatial: bool):
|
||||
if temporal_to_spatial:
|
||||
scatter_dim, gather_dim = 2, 1
|
||||
new_s, new_t = s // get_sequence_parallel_size(), t * get_sequence_parallel_size()
|
||||
else:
|
||||
scatter_dim, gather_dim = 1, 2
|
||||
new_s, new_t = s * get_sequence_parallel_size(), t // get_sequence_parallel_size()
|
||||
|
||||
x = rearrange(x, "b (t s) d -> b t s d", t=t, s=s)
|
||||
x = all_to_all_comm(x, get_sequence_parallel_group(), scatter_dim=scatter_dim, gather_dim=gather_dim)
|
||||
x = rearrange(x, "b t s d -> b (t s) d", t=new_t, s=new_s)
|
||||
return x, new_s, new_t
|
||||
|
||||
|
||||
class STDiT2Config(PretrainedConfig):
|
||||
model_type = "STDiT2"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_size=(None, None, None),
|
||||
input_sq_size=32,
|
||||
in_channels=4,
|
||||
patch_size=(1, 2, 2),
|
||||
hidden_size=1152,
|
||||
depth=28,
|
||||
num_heads=16,
|
||||
mlp_ratio=4.0,
|
||||
class_dropout_prob=0.1,
|
||||
pred_sigma=True,
|
||||
drop_path=0.0,
|
||||
no_temporal_pos_emb=False,
|
||||
caption_channels=4096,
|
||||
model_max_length=120,
|
||||
freeze=None,
|
||||
qk_norm=False,
|
||||
enable_flash_attn=False,
|
||||
enable_layernorm_kernel=False,
|
||||
**kwargs,
|
||||
):
|
||||
self.input_size = input_size
|
||||
self.input_sq_size = input_sq_size
|
||||
self.in_channels = in_channels
|
||||
self.patch_size = patch_size
|
||||
self.hidden_size = hidden_size
|
||||
self.depth = depth
|
||||
self.num_heads = num_heads
|
||||
self.mlp_ratio = mlp_ratio
|
||||
self.class_dropout_prob = class_dropout_prob
|
||||
self.pred_sigma = pred_sigma
|
||||
self.drop_path = drop_path
|
||||
self.no_temporal_pos_emb = no_temporal_pos_emb
|
||||
self.caption_channels = caption_channels
|
||||
self.model_max_length = model_max_length
|
||||
self.freeze = freeze
|
||||
self.qk_norm = qk_norm
|
||||
self.enable_flash_attn = enable_flash_attn
|
||||
self.enable_layernorm_kernel = enable_layernorm_kernel
|
||||
super().__init__(**kwargs)
|
||||
|
||||
|
||||
class STDiT2(PreTrainedModel):
|
||||
config_class = STDiT2Config
|
||||
|
||||
def __init__(self, config):
|
||||
super().__init__(config)
|
||||
self.pred_sigma = config.pred_sigma
|
||||
self.in_channels = config.in_channels
|
||||
self.out_channels = config.in_channels * 2 if config.pred_sigma else config.in_channels
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_heads = config.num_heads
|
||||
self.no_temporal_pos_emb = config.no_temporal_pos_emb
|
||||
self.depth = config.depth
|
||||
self.mlp_ratio = config.mlp_ratio
|
||||
self.enable_flash_attn = config.enable_flash_attn
|
||||
self.enable_layernorm_kernel = config.enable_layernorm_kernel
|
||||
|
||||
# support dynamic input
|
||||
self.patch_size = config.patch_size
|
||||
self.input_size = config.input_size
|
||||
self.input_sq_size = config.input_sq_size
|
||||
self.pos_embed = PositionEmbedding2D(config.hidden_size)
|
||||
|
||||
self.x_embedder = PatchEmbed3D(config.patch_size, config.in_channels, config.hidden_size)
|
||||
self.t_embedder = TimestepEmbedder(config.hidden_size)
|
||||
self.t_block = nn.Sequential(nn.SiLU(), nn.Linear(config.hidden_size, 6 * config.hidden_size, bias=True))
|
||||
self.t_block_temp = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(config.hidden_size, 3 * config.hidden_size, bias=True)
|
||||
) # new
|
||||
self.y_embedder = CaptionEmbedder(
|
||||
in_channels=config.caption_channels,
|
||||
hidden_size=config.hidden_size,
|
||||
uncond_prob=config.class_dropout_prob,
|
||||
act_layer=approx_gelu,
|
||||
token_num=config.model_max_length,
|
||||
)
|
||||
|
||||
drop_path = [x.item() for x in torch.linspace(0, config.drop_path, config.depth)]
|
||||
from rotary_embedding_torch import RotaryEmbedding
|
||||
|
||||
self.rope = RotaryEmbedding(dim=self.hidden_size // self.num_heads, seq_before_head_dim=True) # new
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
STDiT2Block(
|
||||
self.hidden_size,
|
||||
self.num_heads,
|
||||
mlp_ratio=self.mlp_ratio,
|
||||
drop_path=drop_path[i],
|
||||
enable_flash_attn=self.enable_flash_attn,
|
||||
enable_layernorm_kernel=self.enable_layernorm_kernel,
|
||||
rope=self.rope.rotate_queries_or_keys,
|
||||
qk_norm=config.qk_norm,
|
||||
)
|
||||
for i in range(self.depth)
|
||||
]
|
||||
)
|
||||
self.final_layer = T2IFinalLayer(config.hidden_size, np.prod(self.patch_size), self.out_channels)
|
||||
|
||||
# multi_res
|
||||
assert self.hidden_size % 3 == 0, "hidden_size must be divisible by 3"
|
||||
self.csize_embedder = SizeEmbedder(self.hidden_size // 3)
|
||||
self.ar_embedder = SizeEmbedder(self.hidden_size // 3)
|
||||
self.fl_embedder = SizeEmbedder(self.hidden_size) # new
|
||||
self.fps_embedder = SizeEmbedder(self.hidden_size) # new
|
||||
|
||||
# init model
|
||||
self.initialize_weights()
|
||||
self.initialize_temporal()
|
||||
if config.freeze is not None:
|
||||
assert config.freeze in ["not_temporal", "text"]
|
||||
if config.freeze == "not_temporal":
|
||||
self.freeze_not_temporal()
|
||||
elif config.freeze == "text":
|
||||
self.freeze_text()
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
self.gradient_checkpointing = True
|
||||
|
||||
@staticmethod
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
def get_dynamic_size(self, x):
|
||||
_, _, T, H, W = x.size()
|
||||
if T % self.patch_size[0] != 0:
|
||||
T += self.patch_size[0] - T % self.patch_size[0]
|
||||
if H % self.patch_size[1] != 0:
|
||||
H += self.patch_size[1] - H % self.patch_size[1]
|
||||
if W % self.patch_size[2] != 0:
|
||||
W += self.patch_size[2] - W % self.patch_size[2]
|
||||
T = T // self.patch_size[0]
|
||||
H = H // self.patch_size[1]
|
||||
W = W // self.patch_size[2]
|
||||
return (T, H, W)
|
||||
|
||||
def forward(
|
||||
self, x, timestep, y, mask=None, x_mask=None, num_frames=None, height=None, width=None, ar=None, fps=None
|
||||
):
|
||||
"""
|
||||
Forward pass of STDiT.
|
||||
Args:
|
||||
x (torch.Tensor): latent representation of video; of shape [B, C, T, H, W]
|
||||
timestep (torch.Tensor): diffusion time steps; of shape [B]
|
||||
y (torch.Tensor): representation of prompts; of shape [B, 1, N_token, C]
|
||||
mask (torch.Tensor): mask for selecting prompt tokens; of shape [B, N_token]
|
||||
|
||||
Returns:
|
||||
x (torch.Tensor): output latent representation; of shape [B, C, T, H, W]
|
||||
"""
|
||||
B = x.shape[0]
|
||||
dtype = self.x_embedder.proj.weight.dtype
|
||||
x = x.to(dtype)
|
||||
timestep = timestep.to(dtype)
|
||||
y = y.to(dtype)
|
||||
|
||||
# === process data info ===
|
||||
# 1. get dynamic size
|
||||
hw = torch.cat([height[:, None], width[:, None]], dim=1)
|
||||
rs = (height[0].item() * width[0].item()) ** 0.5
|
||||
csize = self.csize_embedder(hw, B)
|
||||
|
||||
# 2. get aspect ratio
|
||||
ar = ar.unsqueeze(1)
|
||||
ar = self.ar_embedder(ar, B)
|
||||
data_info = torch.cat([csize, ar], dim=1)
|
||||
|
||||
# 3. get number of frames
|
||||
fl = num_frames.unsqueeze(1)
|
||||
fps = fps.unsqueeze(1)
|
||||
fl = self.fl_embedder(fl, B)
|
||||
fl = fl + self.fps_embedder(fps, B)
|
||||
|
||||
# === get dynamic shape size ===
|
||||
_, _, Tx, Hx, Wx = x.size()
|
||||
T, H, W = self.get_dynamic_size(x)
|
||||
S = H * W
|
||||
scale = rs / self.input_sq_size
|
||||
base_size = round(S**0.5)
|
||||
pos_emb = self.pos_embed(x, H, W, scale=scale, base_size=base_size)
|
||||
|
||||
# embedding
|
||||
x = self.x_embedder(x) # [B, N, C]
|
||||
x = rearrange(x, "B (T S) C -> B T S C", T=T, S=S)
|
||||
x = x + pos_emb
|
||||
|
||||
if is_sequence_parallelism_enable():
|
||||
x = split_sequence(x, get_sequence_parallel_group(), dim=1)
|
||||
T = T // get_sequence_parallel_size()
|
||||
|
||||
x = rearrange(x, "B T S C -> B (T S) C")
|
||||
|
||||
# prepare adaIN
|
||||
t = self.t_embedder(timestep, dtype=x.dtype) # [B, C]
|
||||
t_spc = t + data_info # [B, C]
|
||||
t_tmp = t + fl # [B, C]
|
||||
t_spc_mlp = self.t_block(t_spc) # [B, 6*C]
|
||||
t_tmp_mlp = self.t_block_temp(t_tmp) # [B, 3*C]
|
||||
if x_mask is not None:
|
||||
t0_timestep = torch.zeros_like(timestep)
|
||||
t0 = self.t_embedder(t0_timestep, dtype=x.dtype)
|
||||
t0_spc = t0 + data_info
|
||||
t0_tmp = t0 + fl
|
||||
t0_spc_mlp = self.t_block(t0_spc)
|
||||
t0_tmp_mlp = self.t_block_temp(t0_tmp)
|
||||
else:
|
||||
t0_spc = None
|
||||
t0_tmp = None
|
||||
t0_spc_mlp = None
|
||||
t0_tmp_mlp = None
|
||||
|
||||
# prepare y
|
||||
y = self.y_embedder(y, self.training) # [B, 1, N_token, C]
|
||||
|
||||
if mask is not None:
|
||||
if mask.shape[0] != y.shape[0]:
|
||||
mask = mask.repeat(y.shape[0] // mask.shape[0], 1)
|
||||
mask = mask.squeeze(1).squeeze(1)
|
||||
y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1])
|
||||
y_lens = mask.sum(dim=1).tolist()
|
||||
else:
|
||||
y_lens = [y.shape[2]] * y.shape[0]
|
||||
y = y.squeeze(1).view(1, -1, x.shape[-1])
|
||||
|
||||
# blocks
|
||||
for _, block in enumerate(self.blocks):
|
||||
if self.gradient_checkpointing:
|
||||
x = torch.utils.checkpoint.checkpoint(
|
||||
self.create_custom_forward(block),
|
||||
x,
|
||||
y,
|
||||
t_spc_mlp,
|
||||
t_tmp_mlp,
|
||||
y_lens,
|
||||
x_mask,
|
||||
t0_spc_mlp,
|
||||
t0_tmp_mlp,
|
||||
T,
|
||||
S,
|
||||
)
|
||||
else:
|
||||
x = block(
|
||||
x,
|
||||
y,
|
||||
t_spc_mlp,
|
||||
t_tmp_mlp,
|
||||
y_lens,
|
||||
x_mask,
|
||||
t0_spc_mlp,
|
||||
t0_tmp_mlp,
|
||||
T,
|
||||
S,
|
||||
)
|
||||
# x.shape: [B, N, C]
|
||||
|
||||
if is_sequence_parallelism_enable():
|
||||
x = gather_sequence(x, get_sequence_parallel_group(), dim=1)
|
||||
T = T * get_sequence_parallel_size()
|
||||
|
||||
# final process
|
||||
x = self.final_layer(x, t, x_mask, t0_spc, T, S) # [B, N, C=T_p * H_p * W_p * C_out]
|
||||
x = self.unpatchify(x, T, H, W, Tx, Hx, Wx) # [B, C_out, T, H, W]
|
||||
|
||||
# cast to float32 for better accuracy
|
||||
x = x.to(torch.float32)
|
||||
return x
|
||||
|
||||
def unpatchify(self, x, N_t, N_h, N_w, R_t, R_h, R_w):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): of shape [B, N, C]
|
||||
|
||||
Return:
|
||||
x (torch.Tensor): of shape [B, C_out, T, H, W]
|
||||
"""
|
||||
|
||||
# N_t, N_h, N_w = [self.input_size[i] // self.patch_size[i] for i in range(3)]
|
||||
T_p, H_p, W_p = self.patch_size
|
||||
x = rearrange(
|
||||
x,
|
||||
"B (N_t N_h N_w) (T_p H_p W_p C_out) -> B C_out (N_t T_p) (N_h H_p) (N_w W_p)",
|
||||
N_t=N_t,
|
||||
N_h=N_h,
|
||||
N_w=N_w,
|
||||
T_p=T_p,
|
||||
H_p=H_p,
|
||||
W_p=W_p,
|
||||
C_out=self.out_channels,
|
||||
)
|
||||
# unpad
|
||||
x = x[:, :, :R_t, :R_h, :R_w]
|
||||
return x
|
||||
|
||||
def unpatchify_old(self, x):
|
||||
c = self.out_channels
|
||||
t, h, w = [self.input_size[i] // self.patch_size[i] for i in range(3)]
|
||||
pt, ph, pw = self.patch_size
|
||||
|
||||
x = x.reshape(shape=(x.shape[0], t, h, w, pt, ph, pw, c))
|
||||
x = rearrange(x, "n t h w r p q c -> n c t r h p w q")
|
||||
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
|
||||
return imgs
|
||||
|
||||
def get_spatial_pos_embed(self, H, W, scale=1.0, base_size=None):
|
||||
pos_embed = get_2d_sincos_pos_embed(
|
||||
self.hidden_size,
|
||||
(H, W),
|
||||
scale=scale,
|
||||
base_size=base_size,
|
||||
)
|
||||
pos_embed = torch.from_numpy(pos_embed).float().unsqueeze(0).requires_grad_(False)
|
||||
return pos_embed
|
||||
|
||||
def freeze_not_temporal(self):
|
||||
for n, p in self.named_parameters():
|
||||
if "attn_temp" not in n:
|
||||
p.requires_grad = False
|
||||
|
||||
def freeze_text(self):
|
||||
for n, p in self.named_parameters():
|
||||
if "cross_attn" in n:
|
||||
p.requires_grad = False
|
||||
|
||||
def initialize_temporal(self):
|
||||
for block in self.blocks:
|
||||
nn.init.constant_(block.attn_temp.proj.weight, 0)
|
||||
nn.init.constant_(block.attn_temp.proj.bias, 0)
|
||||
|
||||
def initialize_weights(self):
|
||||
# Initialize transformer layers:
|
||||
def _basic_init(module):
|
||||
if isinstance(module, nn.Linear):
|
||||
torch.nn.init.xavier_uniform_(module.weight)
|
||||
if module.bias is not None:
|
||||
nn.init.constant_(module.bias, 0)
|
||||
|
||||
self.apply(_basic_init)
|
||||
|
||||
# Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
|
||||
w = self.x_embedder.proj.weight.data
|
||||
nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
|
||||
|
||||
# Initialize timestep embedding MLP:
|
||||
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
|
||||
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
|
||||
nn.init.normal_(self.t_block[1].weight, std=0.02)
|
||||
nn.init.normal_(self.t_block_temp[1].weight, std=0.02)
|
||||
|
||||
# Initialize caption embedding MLP:
|
||||
nn.init.normal_(self.y_embedder.y_proj.fc1.weight, std=0.02)
|
||||
nn.init.normal_(self.y_embedder.y_proj.fc2.weight, std=0.02)
|
||||
|
||||
# Zero-out adaLN modulation layers in PixArt blocks:
|
||||
for block in self.blocks:
|
||||
nn.init.constant_(block.cross_attn.proj.weight, 0)
|
||||
nn.init.constant_(block.cross_attn.proj.bias, 0)
|
||||
|
||||
# Zero-out output layers:
|
||||
nn.init.constant_(self.final_layer.linear.weight, 0)
|
||||
nn.init.constant_(self.final_layer.linear.bias, 0)
|
||||
|
||||
|
||||
def STDiT2_XL_2(from_pretrained=None, **kwargs):
|
||||
if from_pretrained is not None:
|
||||
if os.path.isdir(from_pretrained) or os.path.isfile(from_pretrained):
|
||||
# if it is a directory or a file, we load the checkpoint manually
|
||||
config = STDiT2Config(depth=28, hidden_size=1152, patch_size=(1, 2, 2), num_heads=16, **kwargs)
|
||||
model = STDiT2(config)
|
||||
load_checkpoint(model, from_pretrained)
|
||||
return model
|
||||
else:
|
||||
# otherwise, we load the model from hugging face hub
|
||||
return STDiT2.from_pretrained(from_pretrained, **kwargs)
|
||||
else:
|
||||
# create a new model
|
||||
config = STDiT2Config(depth=28, hidden_size=1152, patch_size=(1, 2, 2), num_heads=16, **kwargs)
|
||||
model = STDiT2(config)
|
||||
return model
|
||||
Executable
+522
@@ -0,0 +1,522 @@
|
||||
# Adapted from OpenSora
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# OpenSora: https://github.com/hpcaitech/Open-Sora
|
||||
# --------------------------------------------------------
|
||||
|
||||
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from timm.models.layers import DropPath
|
||||
from timm.models.vision_transformer import Mlp
|
||||
from transformers import PretrainedConfig, PreTrainedModel
|
||||
|
||||
from opendit.core.comm import (
|
||||
all_to_all_with_pad,
|
||||
gather_sequence,
|
||||
get_spatial_pad,
|
||||
get_temporal_pad,
|
||||
set_spatial_pad,
|
||||
set_temporal_pad,
|
||||
split_sequence,
|
||||
)
|
||||
from opendit.core.pab_mgr import enable_pab, if_broadcast_cross, if_broadcast_spatial, if_broadcast_temporal
|
||||
#from opendit.core.parallel_mgr import enable_sequence_parallel, get_sequence_parallel_group
|
||||
|
||||
def enable_sequence_parallel():
|
||||
return False
|
||||
|
||||
from .modules import (
|
||||
Attention,
|
||||
CaptionEmbedder,
|
||||
MultiHeadCrossAttention,
|
||||
PatchEmbed3D,
|
||||
PositionEmbedding2D,
|
||||
SizeEmbedder,
|
||||
T2IFinalLayer,
|
||||
TimestepEmbedder,
|
||||
approx_gelu,
|
||||
get_layernorm,
|
||||
t2i_modulate,
|
||||
)
|
||||
from .utils import auto_grad_checkpoint, load_checkpoint
|
||||
|
||||
|
||||
class STDiT3Block(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
num_heads,
|
||||
mlp_ratio=4.0,
|
||||
drop_path=0.0,
|
||||
rope=None,
|
||||
qk_norm=False,
|
||||
temporal=False,
|
||||
enable_flash_attn=False,
|
||||
block_idx=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.temporal = temporal
|
||||
self.hidden_size = hidden_size
|
||||
self.enable_flash_attn = enable_flash_attn
|
||||
|
||||
attn_cls = Attention
|
||||
mha_cls = MultiHeadCrossAttention
|
||||
|
||||
self.norm1 = get_layernorm(hidden_size, eps=1e-6, affine=False)
|
||||
self.attn = attn_cls(
|
||||
hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
qk_norm=qk_norm,
|
||||
rope=rope,
|
||||
enable_flash_attn=enable_flash_attn,
|
||||
)
|
||||
self.cross_attn = mha_cls(hidden_size, num_heads)
|
||||
self.norm2 = get_layernorm(hidden_size, eps=1e-6, affine=False)
|
||||
self.mlp = Mlp(
|
||||
in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0
|
||||
)
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(6, hidden_size) / hidden_size**0.5)
|
||||
|
||||
# fast video diffusion
|
||||
self.block_idx = block_idx
|
||||
self.attn_count = 0
|
||||
self.last_attn = None
|
||||
self.cross_count = 0
|
||||
self.last_cross = None
|
||||
|
||||
def t_mask_select(self, x_mask, x, masked_x, T, S):
|
||||
# x: [B, (T, S), C]
|
||||
# mased_x: [B, (T, S), C]
|
||||
# x_mask: [B, T]
|
||||
x = rearrange(x, "B (T S) C -> B T S C", T=T, S=S)
|
||||
masked_x = rearrange(masked_x, "B (T S) C -> B T S C", T=T, S=S)
|
||||
x = torch.where(x_mask[:, :, None, None], x, masked_x)
|
||||
x = rearrange(x, "B T S C -> B (T S) C")
|
||||
return x
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
y,
|
||||
t,
|
||||
mask=None, # text mask
|
||||
x_mask=None, # temporal mask
|
||||
t0=None, # t with timestamp=0
|
||||
T=None, # number of frames
|
||||
S=None, # number of pixel patches
|
||||
timestep=None,
|
||||
):
|
||||
# prepare modulate parameters
|
||||
B, N, C = x.shape
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||
self.scale_shift_table[None] + t.reshape(B, 6, -1)
|
||||
).chunk(6, dim=1)
|
||||
if x_mask is not None:
|
||||
shift_msa_zero, scale_msa_zero, gate_msa_zero, shift_mlp_zero, scale_mlp_zero, gate_mlp_zero = (
|
||||
self.scale_shift_table[None] + t0.reshape(B, 6, -1)
|
||||
).chunk(6, dim=1)
|
||||
|
||||
if self.temporal:
|
||||
broadcast_attn, self.attn_count = if_broadcast_temporal(int(timestep[0]), self.attn_count)
|
||||
else:
|
||||
broadcast_attn, self.attn_count = if_broadcast_spatial(int(timestep[0]), self.attn_count, self.block_idx)
|
||||
|
||||
if broadcast_attn:
|
||||
x_m_s = self.last_attn
|
||||
else:
|
||||
# modulate (attention)
|
||||
x_m = t2i_modulate(self.norm1(x), shift_msa, scale_msa)
|
||||
if x_mask is not None:
|
||||
x_m_zero = t2i_modulate(self.norm1(x), shift_msa_zero, scale_msa_zero)
|
||||
x_m = self.t_mask_select(x_mask, x_m, x_m_zero, T, S)
|
||||
|
||||
# attention
|
||||
if self.temporal:
|
||||
if enable_sequence_parallel():
|
||||
x_m, S, T = self.dynamic_switch(x_m, S, T, to_spatial_shard=True)
|
||||
x_m = rearrange(x_m, "B (T S) C -> (B S) T C", T=T, S=S)
|
||||
x_m = self.attn(x_m)
|
||||
x_m = rearrange(x_m, "(B S) T C -> B (T S) C", T=T, S=S)
|
||||
if enable_sequence_parallel():
|
||||
x_m, S, T = self.dynamic_switch(x_m, S, T, to_spatial_shard=False)
|
||||
else:
|
||||
x_m = rearrange(x_m, "B (T S) C -> (B T) S C", T=T, S=S)
|
||||
x_m = self.attn(x_m)
|
||||
x_m = rearrange(x_m, "(B T) S C -> B (T S) C", T=T, S=S)
|
||||
|
||||
# modulate (attention)
|
||||
x_m_s = gate_msa * x_m
|
||||
if x_mask is not None:
|
||||
x_m_s_zero = gate_msa_zero * x_m
|
||||
x_m_s = self.t_mask_select(x_mask, x_m_s, x_m_s_zero, T, S)
|
||||
|
||||
if enable_pab():
|
||||
self.last_attn = x_m_s
|
||||
|
||||
# residual
|
||||
x = x + self.drop_path(x_m_s)
|
||||
|
||||
# cross attention
|
||||
broadcast_cross, self.cross_count = if_broadcast_cross(int(timestep[0]), self.cross_count)
|
||||
if broadcast_cross:
|
||||
x = x + self.last_cross
|
||||
else:
|
||||
x_cross = self.cross_attn(x, y, mask)
|
||||
if enable_pab():
|
||||
self.last_cross = x_cross
|
||||
x = x + x_cross
|
||||
|
||||
# modulate (MLP)
|
||||
x_m = t2i_modulate(self.norm2(x), shift_mlp, scale_mlp)
|
||||
if x_mask is not None:
|
||||
x_m_zero = t2i_modulate(self.norm2(x), shift_mlp_zero, scale_mlp_zero)
|
||||
x_m = self.t_mask_select(x_mask, x_m, x_m_zero, T, S)
|
||||
|
||||
# MLP
|
||||
x_m = self.mlp(x_m)
|
||||
|
||||
# modulate (MLP)
|
||||
x_m_s = gate_mlp * x_m
|
||||
if x_mask is not None:
|
||||
x_m_s_zero = gate_mlp_zero * x_m
|
||||
x_m_s = self.t_mask_select(x_mask, x_m_s, x_m_s_zero, T, S)
|
||||
|
||||
# residual
|
||||
x = x + self.drop_path(x_m_s)
|
||||
|
||||
return x
|
||||
|
||||
def dynamic_switch(self, x, s, t, to_spatial_shard: bool):
|
||||
if to_spatial_shard:
|
||||
scatter_dim, gather_dim = 2, 1
|
||||
scatter_pad = get_spatial_pad()
|
||||
gather_pad = get_temporal_pad()
|
||||
else:
|
||||
scatter_dim, gather_dim = 1, 2
|
||||
scatter_pad = get_temporal_pad()
|
||||
gather_pad = get_spatial_pad()
|
||||
|
||||
x = rearrange(x, "b (t s) d -> b t s d", t=t, s=s)
|
||||
x = all_to_all_with_pad(
|
||||
x,
|
||||
get_sequence_parallel_group(),
|
||||
scatter_dim=scatter_dim,
|
||||
gather_dim=gather_dim,
|
||||
scatter_pad=scatter_pad,
|
||||
gather_pad=gather_pad,
|
||||
)
|
||||
new_s, new_t = x.shape[2], x.shape[1]
|
||||
x = rearrange(x, "b t s d -> b (t s) d")
|
||||
return x, new_s, new_t
|
||||
|
||||
|
||||
class STDiT3Config(PretrainedConfig):
|
||||
model_type = "STDiT3"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_size=(None, None, None),
|
||||
input_sq_size=512,
|
||||
in_channels=4,
|
||||
patch_size=(1, 2, 2),
|
||||
hidden_size=1152,
|
||||
depth=28,
|
||||
num_heads=16,
|
||||
mlp_ratio=4.0,
|
||||
class_dropout_prob=0.1,
|
||||
pred_sigma=True,
|
||||
drop_path=0.0,
|
||||
caption_channels=4096,
|
||||
model_max_length=300,
|
||||
qk_norm=True,
|
||||
enable_flash_attn=False,
|
||||
only_train_temporal=False,
|
||||
freeze_y_embedder=False,
|
||||
skip_y_embedder=False,
|
||||
**kwargs,
|
||||
):
|
||||
self.input_size = input_size
|
||||
self.input_sq_size = input_sq_size
|
||||
self.in_channels = in_channels
|
||||
self.patch_size = patch_size
|
||||
self.hidden_size = hidden_size
|
||||
self.depth = depth
|
||||
self.num_heads = num_heads
|
||||
self.mlp_ratio = mlp_ratio
|
||||
self.class_dropout_prob = class_dropout_prob
|
||||
self.pred_sigma = pred_sigma
|
||||
self.drop_path = drop_path
|
||||
self.caption_channels = caption_channels
|
||||
self.model_max_length = model_max_length
|
||||
self.qk_norm = qk_norm
|
||||
self.enable_flash_attn = enable_flash_attn
|
||||
self.only_train_temporal = only_train_temporal
|
||||
self.freeze_y_embedder = freeze_y_embedder
|
||||
self.skip_y_embedder = skip_y_embedder
|
||||
super().__init__(**kwargs)
|
||||
|
||||
|
||||
class STDiT3(PreTrainedModel):
|
||||
config_class = STDiT3Config
|
||||
|
||||
def __init__(self, config):
|
||||
super().__init__(config)
|
||||
self.pred_sigma = config.pred_sigma
|
||||
self.in_channels = config.in_channels
|
||||
self.out_channels = config.in_channels * 2 if config.pred_sigma else config.in_channels
|
||||
|
||||
# model size related
|
||||
self.depth = config.depth
|
||||
self.mlp_ratio = config.mlp_ratio
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_heads = config.num_heads
|
||||
|
||||
# computation related
|
||||
self.drop_path = config.drop_path
|
||||
self.enable_flash_attn = config.enable_flash_attn
|
||||
|
||||
# input size related
|
||||
self.patch_size = config.patch_size
|
||||
self.input_sq_size = config.input_sq_size
|
||||
self.pos_embed = PositionEmbedding2D(config.hidden_size)
|
||||
|
||||
from rotary_embedding_torch import RotaryEmbedding
|
||||
|
||||
self.rope = RotaryEmbedding(dim=self.hidden_size // self.num_heads)
|
||||
|
||||
# embedding
|
||||
self.x_embedder = PatchEmbed3D(config.patch_size, config.in_channels, config.hidden_size)
|
||||
self.t_embedder = TimestepEmbedder(config.hidden_size)
|
||||
self.fps_embedder = SizeEmbedder(self.hidden_size)
|
||||
self.t_block = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(config.hidden_size, 6 * config.hidden_size, bias=True),
|
||||
)
|
||||
self.y_embedder = CaptionEmbedder(
|
||||
in_channels=config.caption_channels,
|
||||
hidden_size=config.hidden_size,
|
||||
uncond_prob=config.class_dropout_prob,
|
||||
act_layer=approx_gelu,
|
||||
token_num=config.model_max_length,
|
||||
)
|
||||
|
||||
# spatial blocks
|
||||
drop_path = [x.item() for x in torch.linspace(0, self.drop_path, config.depth)]
|
||||
self.spatial_blocks = nn.ModuleList(
|
||||
[
|
||||
STDiT3Block(
|
||||
hidden_size=config.hidden_size,
|
||||
num_heads=config.num_heads,
|
||||
mlp_ratio=config.mlp_ratio,
|
||||
drop_path=drop_path[i],
|
||||
qk_norm=config.qk_norm,
|
||||
enable_flash_attn=config.enable_flash_attn,
|
||||
block_idx=i,
|
||||
)
|
||||
for i in range(config.depth)
|
||||
]
|
||||
)
|
||||
|
||||
# temporal blocks
|
||||
drop_path = [x.item() for x in torch.linspace(0, self.drop_path, config.depth)]
|
||||
self.temporal_blocks = nn.ModuleList(
|
||||
[
|
||||
STDiT3Block(
|
||||
hidden_size=config.hidden_size,
|
||||
num_heads=config.num_heads,
|
||||
mlp_ratio=config.mlp_ratio,
|
||||
drop_path=drop_path[i],
|
||||
qk_norm=config.qk_norm,
|
||||
enable_flash_attn=config.enable_flash_attn,
|
||||
# temporal
|
||||
temporal=True,
|
||||
rope=self.rope.rotate_queries_or_keys,
|
||||
block_idx=i,
|
||||
)
|
||||
for i in range(config.depth)
|
||||
]
|
||||
)
|
||||
|
||||
# final layer
|
||||
self.final_layer = T2IFinalLayer(config.hidden_size, np.prod(self.patch_size), self.out_channels)
|
||||
|
||||
self.initialize_weights()
|
||||
if config.only_train_temporal:
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
for block in self.temporal_blocks:
|
||||
for param in block.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
if config.freeze_y_embedder:
|
||||
for param in self.y_embedder.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def initialize_weights(self):
|
||||
# Initialize transformer layers:
|
||||
def _basic_init(module):
|
||||
if isinstance(module, nn.Linear):
|
||||
torch.nn.init.xavier_uniform_(module.weight)
|
||||
if module.bias is not None:
|
||||
nn.init.constant_(module.bias, 0)
|
||||
|
||||
self.apply(_basic_init)
|
||||
|
||||
# Initialize fps_embedder
|
||||
nn.init.normal_(self.fps_embedder.mlp[0].weight, std=0.02)
|
||||
nn.init.constant_(self.fps_embedder.mlp[0].bias, 0)
|
||||
nn.init.constant_(self.fps_embedder.mlp[2].weight, 0)
|
||||
nn.init.constant_(self.fps_embedder.mlp[2].bias, 0)
|
||||
|
||||
# Initialize timporal blocks
|
||||
for block in self.temporal_blocks:
|
||||
nn.init.constant_(block.attn.proj.weight, 0)
|
||||
nn.init.constant_(block.cross_attn.proj.weight, 0)
|
||||
nn.init.constant_(block.mlp.fc2.weight, 0)
|
||||
|
||||
def get_dynamic_size(self, x):
|
||||
_, _, T, H, W = x.size()
|
||||
if T % self.patch_size[0] != 0:
|
||||
T += self.patch_size[0] - T % self.patch_size[0]
|
||||
if H % self.patch_size[1] != 0:
|
||||
H += self.patch_size[1] - H % self.patch_size[1]
|
||||
if W % self.patch_size[2] != 0:
|
||||
W += self.patch_size[2] - W % self.patch_size[2]
|
||||
T = T // self.patch_size[0]
|
||||
H = H // self.patch_size[1]
|
||||
W = W // self.patch_size[2]
|
||||
return (T, H, W)
|
||||
|
||||
def encode_text(self, y, mask=None):
|
||||
y = self.y_embedder(y, self.training) # [B, 1, N_token, C]
|
||||
if mask is not None:
|
||||
if mask.shape[0] != y.shape[0]:
|
||||
mask = mask.repeat(y.shape[0] // mask.shape[0], 1)
|
||||
mask = mask.squeeze(1).squeeze(1)
|
||||
y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, self.hidden_size)
|
||||
y_lens = mask.sum(dim=1).tolist()
|
||||
else:
|
||||
y_lens = [y.shape[2]] * y.shape[0]
|
||||
y = y.squeeze(1).view(1, -1, self.hidden_size)
|
||||
return y, y_lens
|
||||
|
||||
def forward(self, x, timestep, y, mask=None, x_mask=None, fps=None, height=None, width=None, **kwargs):
|
||||
dtype = self.x_embedder.proj.weight.dtype
|
||||
B = x.size(0)
|
||||
x = x.to(dtype)
|
||||
timestep = timestep.to(dtype)
|
||||
y = y.to(dtype)
|
||||
|
||||
# === get pos embed ===
|
||||
_, _, Tx, Hx, Wx = x.size()
|
||||
T, H, W = self.get_dynamic_size(x)
|
||||
S = H * W
|
||||
base_size = round(S**0.5)
|
||||
resolution_sq = (height[0].item() * width[0].item()) ** 0.5
|
||||
scale = resolution_sq / self.input_sq_size
|
||||
pos_emb = self.pos_embed(x, H, W, scale=scale, base_size=base_size)
|
||||
|
||||
# === get timestep embed ===
|
||||
t = self.t_embedder(timestep, dtype=x.dtype) # [B, C]
|
||||
fps = self.fps_embedder(fps.unsqueeze(1), B)
|
||||
t = t + fps
|
||||
t_mlp = self.t_block(t)
|
||||
t0 = t0_mlp = None
|
||||
if x_mask is not None:
|
||||
t0_timestep = torch.zeros_like(timestep)
|
||||
t0 = self.t_embedder(t0_timestep, dtype=x.dtype)
|
||||
t0 = t0 + fps
|
||||
t0_mlp = self.t_block(t0)
|
||||
|
||||
# === get y embed ===
|
||||
if self.config.skip_y_embedder:
|
||||
y_lens = mask
|
||||
if isinstance(y_lens, torch.Tensor):
|
||||
y_lens = y_lens.long().tolist()
|
||||
else:
|
||||
y, y_lens = self.encode_text(y, mask)
|
||||
|
||||
# === get x embed ===
|
||||
x = self.x_embedder(x) # [B, N, C]
|
||||
x = rearrange(x, "B (T S) C -> B T S C", T=T, S=S)
|
||||
x = x + pos_emb
|
||||
|
||||
# shard over the sequence dim if sp is enabled
|
||||
if enable_sequence_parallel():
|
||||
set_temporal_pad(T)
|
||||
set_spatial_pad(S)
|
||||
x = split_sequence(x, get_sequence_parallel_group(), dim=1, grad_scale="down", pad=get_temporal_pad())
|
||||
T = x.shape[1]
|
||||
x_mask_org = x_mask
|
||||
x_mask = split_sequence(
|
||||
x_mask, get_sequence_parallel_group(), dim=1, grad_scale="down", pad=get_temporal_pad()
|
||||
)
|
||||
|
||||
x = rearrange(x, "B T S C -> B (T S) C", T=T, S=S)
|
||||
|
||||
# === blocks ===
|
||||
for spatial_block, temporal_block in zip(self.spatial_blocks, self.temporal_blocks):
|
||||
x = auto_grad_checkpoint(spatial_block, x, y, t_mlp, y_lens, x_mask, t0_mlp, T, S, timestep)
|
||||
x = auto_grad_checkpoint(temporal_block, x, y, t_mlp, y_lens, x_mask, t0_mlp, T, S, timestep)
|
||||
|
||||
if enable_sequence_parallel():
|
||||
x = rearrange(x, "B (T S) C -> B T S C", T=T, S=S)
|
||||
x = gather_sequence(x, get_sequence_parallel_group(), dim=1, grad_scale="up", pad=get_temporal_pad())
|
||||
T, S = x.shape[1], x.shape[2]
|
||||
x = rearrange(x, "B T S C -> B (T S) C", T=T, S=S)
|
||||
x_mask = x_mask_org
|
||||
|
||||
# === final layer ===
|
||||
x = self.final_layer(x, t, x_mask, t0, T, S)
|
||||
x = self.unpatchify(x, T, H, W, Tx, Hx, Wx)
|
||||
|
||||
# cast to float32 for better accuracy
|
||||
x = x.to(torch.float32)
|
||||
return x
|
||||
|
||||
def unpatchify(self, x, N_t, N_h, N_w, R_t, R_h, R_w):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): of shape [B, N, C]
|
||||
|
||||
Return:
|
||||
x (torch.Tensor): of shape [B, C_out, T, H, W]
|
||||
"""
|
||||
|
||||
# N_t, N_h, N_w = [self.input_size[i] // self.patch_size[i] for i in range(3)]
|
||||
T_p, H_p, W_p = self.patch_size
|
||||
x = rearrange(
|
||||
x,
|
||||
"B (N_t N_h N_w) (T_p H_p W_p C_out) -> B C_out (N_t T_p) (N_h H_p) (N_w W_p)",
|
||||
N_t=N_t,
|
||||
N_h=N_h,
|
||||
N_w=N_w,
|
||||
T_p=T_p,
|
||||
H_p=H_p,
|
||||
W_p=W_p,
|
||||
C_out=self.out_channels,
|
||||
)
|
||||
# unpad
|
||||
x = x[:, :, :R_t, :R_h, :R_w]
|
||||
return x
|
||||
|
||||
|
||||
def STDiT3_XL_2(from_pretrained=None, **kwargs):
|
||||
#if from_pretrained is not None and not os.path.isdir(from_pretrained):
|
||||
model = STDiT3.from_pretrained(from_pretrained, **kwargs)
|
||||
# else:
|
||||
# config = STDiT3Config(depth=28, hidden_size=1152, patch_size=(1, 2, 2), num_heads=16, **kwargs)
|
||||
# model = STDiT3(config)
|
||||
# if from_pretrained is not None:
|
||||
# load_checkpoint(model, from_pretrained)
|
||||
return model
|
||||
Executable
+321
@@ -0,0 +1,321 @@
|
||||
# Adapted from OpenSora
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# OpenSora: https://github.com/hpcaitech/Open-Sora
|
||||
# --------------------------------------------------------
|
||||
|
||||
import html
|
||||
import re
|
||||
|
||||
import ftfy
|
||||
import torch
|
||||
from transformers import AutoTokenizer, T5EncoderModel
|
||||
|
||||
|
||||
class T5Embedder:
|
||||
available_models = ["DeepFloyd/t5-v1_1-xxl"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
device,
|
||||
from_pretrained=None,
|
||||
*,
|
||||
cache_dir=None,
|
||||
hf_token=None,
|
||||
use_text_preprocessing=True,
|
||||
t5_model_kwargs=None,
|
||||
torch_dtype=None,
|
||||
use_offload_folder=None,
|
||||
model_max_length=120,
|
||||
local_files_only=False,
|
||||
):
|
||||
self.device = torch.device(device)
|
||||
self.torch_dtype = torch_dtype or torch.bfloat16
|
||||
self.cache_dir = cache_dir
|
||||
|
||||
if t5_model_kwargs is None:
|
||||
t5_model_kwargs = {
|
||||
"low_cpu_mem_usage": True,
|
||||
"torch_dtype": self.torch_dtype,
|
||||
}
|
||||
|
||||
if use_offload_folder is not None:
|
||||
t5_model_kwargs["offload_folder"] = use_offload_folder
|
||||
t5_model_kwargs["device_map"] = {
|
||||
"shared": self.device,
|
||||
"encoder.embed_tokens": self.device,
|
||||
"encoder.block.0": self.device,
|
||||
"encoder.block.1": self.device,
|
||||
"encoder.block.2": self.device,
|
||||
"encoder.block.3": self.device,
|
||||
"encoder.block.4": self.device,
|
||||
"encoder.block.5": self.device,
|
||||
"encoder.block.6": self.device,
|
||||
"encoder.block.7": self.device,
|
||||
"encoder.block.8": self.device,
|
||||
"encoder.block.9": self.device,
|
||||
"encoder.block.10": self.device,
|
||||
"encoder.block.11": self.device,
|
||||
"encoder.block.12": "disk",
|
||||
"encoder.block.13": "disk",
|
||||
"encoder.block.14": "disk",
|
||||
"encoder.block.15": "disk",
|
||||
"encoder.block.16": "disk",
|
||||
"encoder.block.17": "disk",
|
||||
"encoder.block.18": "disk",
|
||||
"encoder.block.19": "disk",
|
||||
"encoder.block.20": "disk",
|
||||
"encoder.block.21": "disk",
|
||||
"encoder.block.22": "disk",
|
||||
"encoder.block.23": "disk",
|
||||
"encoder.final_layer_norm": "disk",
|
||||
"encoder.dropout": "disk",
|
||||
}
|
||||
else:
|
||||
t5_model_kwargs["device_map"] = {
|
||||
"shared": self.device,
|
||||
"encoder": self.device,
|
||||
}
|
||||
|
||||
self.use_text_preprocessing = use_text_preprocessing
|
||||
self.hf_token = hf_token
|
||||
|
||||
#assert from_pretrained in self.available_models
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(
|
||||
from_pretrained,
|
||||
cache_dir=cache_dir,
|
||||
local_files_only=local_files_only,
|
||||
)
|
||||
self.model = T5EncoderModel.from_pretrained(
|
||||
from_pretrained,
|
||||
cache_dir=cache_dir,
|
||||
local_files_only=local_files_only,
|
||||
**t5_model_kwargs,
|
||||
).eval()
|
||||
self.model_max_length = model_max_length
|
||||
|
||||
def get_text_embeddings(self, texts):
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
texts,
|
||||
max_length=self.model_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
input_ids = text_tokens_and_mask["input_ids"].to(self.device)
|
||||
attention_mask = text_tokens_and_mask["attention_mask"].to(self.device)
|
||||
with torch.no_grad():
|
||||
text_encoder_embs = self.model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
)["last_hidden_state"].detach()
|
||||
return text_encoder_embs, attention_mask
|
||||
|
||||
|
||||
class T5Encoder:
|
||||
def __init__(
|
||||
self,
|
||||
from_pretrained=None,
|
||||
model_max_length=120,
|
||||
device="cuda",
|
||||
dtype=torch.float,
|
||||
cache_dir=None,
|
||||
shardformer=False,
|
||||
local_files_only=False,
|
||||
):
|
||||
assert from_pretrained is not None, "Please specify the path to the T5 model"
|
||||
|
||||
self.t5 = T5Embedder(
|
||||
device=device,
|
||||
torch_dtype=dtype,
|
||||
from_pretrained=from_pretrained,
|
||||
cache_dir=cache_dir,
|
||||
model_max_length=model_max_length,
|
||||
local_files_only=local_files_only,
|
||||
)
|
||||
self.t5.model.to(dtype=dtype)
|
||||
self.y_embedder = None
|
||||
|
||||
self.model_max_length = model_max_length
|
||||
self.output_dim = self.t5.model.config.d_model
|
||||
self.dtype = dtype
|
||||
|
||||
if shardformer:
|
||||
self.shardformer_t5()
|
||||
|
||||
def shardformer_t5(self):
|
||||
from colossalai.shardformer import ShardConfig, ShardFormer
|
||||
|
||||
from opendit.core.shardformer.t5.policy import T5EncoderPolicy
|
||||
from opendit.utils.utils import requires_grad
|
||||
|
||||
shard_config = ShardConfig(
|
||||
tensor_parallel_process_group=None,
|
||||
pipeline_stage_manager=None,
|
||||
enable_tensor_parallelism=False,
|
||||
enable_fused_normalization=False,
|
||||
enable_flash_attention=False,
|
||||
enable_jit_fused=True,
|
||||
enable_sequence_parallelism=False,
|
||||
enable_sequence_overlap=False,
|
||||
)
|
||||
shard_former = ShardFormer(shard_config=shard_config)
|
||||
optim_model, _ = shard_former.optimize(self.t5.model, policy=T5EncoderPolicy())
|
||||
self.t5.model = optim_model.to(self.dtype)
|
||||
|
||||
# ensure the weights are frozen
|
||||
requires_grad(self.t5.model, False)
|
||||
|
||||
def encode(self, text):
|
||||
caption_embs, emb_masks = self.t5.get_text_embeddings(text)
|
||||
caption_embs = caption_embs[:, None]
|
||||
return dict(y=caption_embs, mask=emb_masks)
|
||||
|
||||
def null(self, n):
|
||||
null_y = self.y_embedder.y_embedding[None].repeat(n, 1, 1)[:, None]
|
||||
return null_y
|
||||
|
||||
|
||||
def basic_clean(text):
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
|
||||
BAD_PUNCT_REGEX = re.compile(
|
||||
r"[" + "#®•©™&@·º½¾¿¡§~" + "\)" + "\(" + "\]" + "\[" + "\}" + "\{" + "\|" + "\\" + "\/" + "\*" + r"]{1,}"
|
||||
) # noqa
|
||||
|
||||
|
||||
def clean_caption(caption):
|
||||
import urllib.parse as ul
|
||||
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
caption = str(caption)
|
||||
caption = ul.unquote_plus(caption)
|
||||
caption = caption.strip().lower()
|
||||
caption = re.sub("<person>", "person", caption)
|
||||
# urls:
|
||||
caption = re.sub(
|
||||
r"\b((?:https?:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))", # noqa
|
||||
"",
|
||||
caption,
|
||||
) # regex for urls
|
||||
caption = re.sub(
|
||||
r"\b((?:www:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))", # noqa
|
||||
"",
|
||||
caption,
|
||||
) # regex for urls
|
||||
# html:
|
||||
caption = BeautifulSoup(caption, features="html.parser").text
|
||||
|
||||
# @<nickname>
|
||||
caption = re.sub(r"@[\w\d]+\b", "", caption)
|
||||
|
||||
# 31C0—31EF CJK Strokes
|
||||
# 31F0—31FF Katakana Phonetic Extensions
|
||||
# 3200—32FF Enclosed CJK Letters and Months
|
||||
# 3300—33FF CJK Compatibility
|
||||
# 3400—4DBF CJK Unified Ideographs Extension A
|
||||
# 4DC0—4DFF Yijing Hexagram Symbols
|
||||
# 4E00—9FFF CJK Unified Ideographs
|
||||
caption = re.sub(r"[\u31c0-\u31ef]+", "", caption)
|
||||
caption = re.sub(r"[\u31f0-\u31ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3200-\u32ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3300-\u33ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3400-\u4dbf]+", "", caption)
|
||||
caption = re.sub(r"[\u4dc0-\u4dff]+", "", caption)
|
||||
caption = re.sub(r"[\u4e00-\u9fff]+", "", caption)
|
||||
#######################################################
|
||||
|
||||
# все виды тире / all types of dash --> "-"
|
||||
caption = re.sub(
|
||||
r"[\u002D\u058A\u05BE\u1400\u1806\u2010-\u2015\u2E17\u2E1A\u2E3A\u2E3B\u2E40\u301C\u3030\u30A0\uFE31\uFE32\uFE58\uFE63\uFF0D]+", # noqa
|
||||
"-",
|
||||
caption,
|
||||
)
|
||||
|
||||
# кавычки к одному стандарту
|
||||
caption = re.sub(r"[`´«»“”¨]", '"', caption)
|
||||
caption = re.sub(r"[‘’]", "'", caption)
|
||||
|
||||
# "
|
||||
caption = re.sub(r""?", "", caption)
|
||||
# &
|
||||
caption = re.sub(r"&", "", caption)
|
||||
|
||||
# ip adresses:
|
||||
caption = re.sub(r"\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}", " ", caption)
|
||||
|
||||
# article ids:
|
||||
caption = re.sub(r"\d:\d\d\s+$", "", caption)
|
||||
|
||||
# \n
|
||||
caption = re.sub(r"\\n", " ", caption)
|
||||
|
||||
# "#123"
|
||||
caption = re.sub(r"#\d{1,3}\b", "", caption)
|
||||
# "#12345.."
|
||||
caption = re.sub(r"#\d{5,}\b", "", caption)
|
||||
# "123456.."
|
||||
caption = re.sub(r"\b\d{6,}\b", "", caption)
|
||||
# filenames:
|
||||
caption = re.sub(r"[\S]+\.(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)", "", caption)
|
||||
|
||||
#
|
||||
caption = re.sub(r"[\"\']{2,}", r'"', caption) # """AUSVERKAUFT"""
|
||||
caption = re.sub(r"[\.]{2,}", r" ", caption) # """AUSVERKAUFT"""
|
||||
|
||||
caption = re.sub(BAD_PUNCT_REGEX, r" ", caption) # ***AUSVERKAUFT***, #AUSVERKAUFT
|
||||
caption = re.sub(r"\s+\.\s+", r" ", caption) # " . "
|
||||
|
||||
# this-is-my-cute-cat / this_is_my_cute_cat
|
||||
regex2 = re.compile(r"(?:\-|\_)")
|
||||
if len(re.findall(regex2, caption)) > 3:
|
||||
caption = re.sub(regex2, " ", caption)
|
||||
|
||||
caption = basic_clean(caption)
|
||||
|
||||
caption = re.sub(r"\b[a-zA-Z]{1,3}\d{3,15}\b", "", caption) # jc6640
|
||||
caption = re.sub(r"\b[a-zA-Z]+\d+[a-zA-Z]+\b", "", caption) # jc6640vc
|
||||
caption = re.sub(r"\b\d+[a-zA-Z]+\d+\b", "", caption) # 6640vc231
|
||||
|
||||
caption = re.sub(r"(worldwide\s+)?(free\s+)?shipping", "", caption)
|
||||
caption = re.sub(r"(free\s)?download(\sfree)?", "", caption)
|
||||
caption = re.sub(r"\bclick\b\s(?:for|on)\s\w+", "", caption)
|
||||
caption = re.sub(r"\b(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)(\simage[s]?)?", "", caption)
|
||||
caption = re.sub(r"\bpage\s+\d+\b", "", caption)
|
||||
|
||||
caption = re.sub(r"\b\d*[a-zA-Z]+\d+[a-zA-Z]+\d+[a-zA-Z\d]*\b", r" ", caption) # j2d1a2a...
|
||||
|
||||
caption = re.sub(r"\b\d+\.?\d*[xх×]\d+\.?\d*\b", "", caption)
|
||||
|
||||
caption = re.sub(r"\b\s+\:\s+", r": ", caption)
|
||||
caption = re.sub(r"(\D[,\./])\b", r"\1 ", caption)
|
||||
caption = re.sub(r"\s+", " ", caption)
|
||||
|
||||
caption.strip()
|
||||
|
||||
caption = re.sub(r"^[\"\']([\w\W]+)[\"\']$", r"\1", caption)
|
||||
caption = re.sub(r"^[\'\_,\-\:;]", r"", caption)
|
||||
caption = re.sub(r"[\'\_,\-\:\-\+]$", r"", caption)
|
||||
caption = re.sub(r"^\.\S+$", "", caption)
|
||||
|
||||
return caption.strip()
|
||||
|
||||
|
||||
def text_preprocessing(text, use_text_preprocessing: bool = True):
|
||||
if use_text_preprocessing:
|
||||
# The exact text cleaning as was in the training stage:
|
||||
text = clean_caption(text)
|
||||
text = clean_caption(text)
|
||||
return text
|
||||
else:
|
||||
return text.lower().strip()
|
||||
Executable
+179
@@ -0,0 +1,179 @@
|
||||
# Adapted from OpenSora
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# OpenSora: https://github.com/hpcaitech/Open-Sora
|
||||
# --------------------------------------------------------
|
||||
|
||||
import os
|
||||
from collections.abc import Iterable
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
#from colossalai.checkpoint_io import GeneralCheckpointIO
|
||||
from torch.utils.checkpoint import checkpoint, checkpoint_sequential
|
||||
from torchvision.datasets.utils import download_url
|
||||
|
||||
from opendit.utils.utils import get_logger
|
||||
|
||||
hf_endpoint = os.environ.get("HF_ENDPOINT")
|
||||
if hf_endpoint is None:
|
||||
hf_endpoint = "https://huggingface.co"
|
||||
|
||||
pretrained_models = {
|
||||
"DiT-XL-2-512x512.pt": "https://dl.fbaipublicfiles.com/DiT/models/DiT-XL-2-512x512.pt",
|
||||
"DiT-XL-2-256x256.pt": "https://dl.fbaipublicfiles.com/DiT/models/DiT-XL-2-256x256.pt",
|
||||
"Latte-XL-2-256x256-ucf101.pt": hf_endpoint + "/maxin-cn/Latte/resolve/main/ucf101.pt",
|
||||
"PixArt-XL-2-256x256.pth": hf_endpoint + "/PixArt-alpha/PixArt-alpha/resolve/main/PixArt-XL-2-256x256.pth",
|
||||
"PixArt-XL-2-SAM-256x256.pth": hf_endpoint + "/PixArt-alpha/PixArt-alpha/resolve/main/PixArt-XL-2-SAM-256x256.pth",
|
||||
"PixArt-XL-2-512x512.pth": hf_endpoint + "/PixArt-alpha/PixArt-alpha/resolve/main/PixArt-XL-2-512x512.pth",
|
||||
"PixArt-XL-2-1024-MS.pth": hf_endpoint + "/PixArt-alpha/PixArt-alpha/resolve/main/PixArt-XL-2-1024-MS.pth",
|
||||
"OpenSora-v1-16x256x256.pth": hf_endpoint + "/hpcai-tech/Open-Sora/resolve/main/OpenSora-v1-16x256x256.pth",
|
||||
"OpenSora-v1-HQ-16x256x256.pth": hf_endpoint + "/hpcai-tech/Open-Sora/resolve/main/OpenSora-v1-HQ-16x256x256.pth",
|
||||
"OpenSora-v1-HQ-16x512x512.pth": hf_endpoint + "/hpcai-tech/Open-Sora/resolve/main/OpenSora-v1-HQ-16x512x512.pth",
|
||||
"PixArt-Sigma-XL-2-256x256.pth": hf_endpoint
|
||||
+ "/PixArt-alpha/PixArt-Sigma/resolve/main/PixArt-Sigma-XL-2-256x256.pth",
|
||||
"PixArt-Sigma-XL-2-512-MS.pth": hf_endpoint
|
||||
+ "/PixArt-alpha/PixArt-Sigma/resolve/main/PixArt-Sigma-XL-2-512-MS.pth",
|
||||
"PixArt-Sigma-XL-2-1024-MS.pth": hf_endpoint
|
||||
+ "/PixArt-alpha/PixArt-Sigma/resolve/main/PixArt-Sigma-XL-2-1024-MS.pth",
|
||||
"PixArt-Sigma-XL-2-2K-MS.pth": hf_endpoint + "/PixArt-alpha/PixArt-Sigma/resolve/main/PixArt-Sigma-XL-2-2K-MS.pth",
|
||||
}
|
||||
|
||||
|
||||
def load_from_sharded_state_dict(model, ckpt_path, model_name="model", strict=False):
|
||||
ckpt_io = GeneralCheckpointIO()
|
||||
ckpt_io.load_model(model, os.path.join(ckpt_path, model_name), strict=strict)
|
||||
|
||||
|
||||
def reparameter(ckpt, name=None, model=None):
|
||||
model_name = name
|
||||
name = os.path.basename(name)
|
||||
if not dist.is_initialized() or dist.get_rank() == 0:
|
||||
get_logger().info("loading pretrained model: %s", model_name)
|
||||
if name in ["DiT-XL-2-512x512.pt", "DiT-XL-2-256x256.pt"]:
|
||||
ckpt["x_embedder.proj.weight"] = ckpt["x_embedder.proj.weight"].unsqueeze(2)
|
||||
del ckpt["pos_embed"]
|
||||
if name in ["Latte-XL-2-256x256-ucf101.pt"]:
|
||||
ckpt = ckpt["ema"]
|
||||
ckpt["x_embedder.proj.weight"] = ckpt["x_embedder.proj.weight"].unsqueeze(2)
|
||||
del ckpt["pos_embed"]
|
||||
del ckpt["temp_embed"]
|
||||
if name in [
|
||||
"PixArt-XL-2-256x256.pth",
|
||||
"PixArt-XL-2-SAM-256x256.pth",
|
||||
"PixArt-XL-2-512x512.pth",
|
||||
"PixArt-XL-2-1024-MS.pth",
|
||||
"PixArt-Sigma-XL-2-256x256.pth",
|
||||
"PixArt-Sigma-XL-2-512-MS.pth",
|
||||
"PixArt-Sigma-XL-2-1024-MS.pth",
|
||||
"PixArt-Sigma-XL-2-2K-MS.pth",
|
||||
]:
|
||||
ckpt = ckpt["state_dict"]
|
||||
ckpt["x_embedder.proj.weight"] = ckpt["x_embedder.proj.weight"].unsqueeze(2)
|
||||
if "pos_embed" in ckpt:
|
||||
del ckpt["pos_embed"]
|
||||
|
||||
if name in [
|
||||
"PixArt-1B-2.pth",
|
||||
]:
|
||||
ckpt = ckpt["state_dict"]
|
||||
if "pos_embed" in ckpt:
|
||||
del ckpt["pos_embed"]
|
||||
|
||||
# no need pos_embed
|
||||
if "pos_embed_temporal" in ckpt:
|
||||
del ckpt["pos_embed_temporal"]
|
||||
if "pos_embed" in ckpt:
|
||||
del ckpt["pos_embed"]
|
||||
# different text length
|
||||
if "y_embedder.y_embedding" in ckpt:
|
||||
if ckpt["y_embedder.y_embedding"].shape[0] < model.y_embedder.y_embedding.shape[0]:
|
||||
get_logger().info(
|
||||
"Extend y_embedding from %s to %s",
|
||||
ckpt["y_embedder.y_embedding"].shape[0],
|
||||
model.y_embedder.y_embedding.shape[0],
|
||||
)
|
||||
additional_length = model.y_embedder.y_embedding.shape[0] - ckpt["y_embedder.y_embedding"].shape[0]
|
||||
new_y_embedding = torch.zeros(additional_length, model.y_embedder.y_embedding.shape[1])
|
||||
new_y_embedding[:] = ckpt["y_embedder.y_embedding"][-1]
|
||||
ckpt["y_embedder.y_embedding"] = torch.cat([ckpt["y_embedder.y_embedding"], new_y_embedding], dim=0)
|
||||
elif ckpt["y_embedder.y_embedding"].shape[0] > model.y_embedder.y_embedding.shape[0]:
|
||||
get_logger().info(
|
||||
"Shrink y_embedding from %s to %s",
|
||||
ckpt["y_embedder.y_embedding"].shape[0],
|
||||
model.y_embedder.y_embedding.shape[0],
|
||||
)
|
||||
ckpt["y_embedder.y_embedding"] = ckpt["y_embedder.y_embedding"][: model.y_embedder.y_embedding.shape[0]]
|
||||
# stdit3 special case
|
||||
if type(model).__name__ == "STDiT3" and "PixArt-Sigma" in name:
|
||||
ckpt_keys = list(ckpt.keys())
|
||||
for key in ckpt_keys:
|
||||
if "blocks." in key:
|
||||
ckpt[key.replace("blocks.", "spatial_blocks.")] = ckpt[key]
|
||||
del ckpt[key]
|
||||
|
||||
return ckpt
|
||||
|
||||
|
||||
def find_model(model_name, model=None):
|
||||
"""
|
||||
Finds a pre-trained DiT model, downloading it if necessary. Alternatively, loads a model from a local path.
|
||||
"""
|
||||
if model_name in pretrained_models: # Find/download our pre-trained DiT checkpoints
|
||||
model_ckpt = download_model(model_name)
|
||||
model_ckpt = reparameter(model_ckpt, model_name, model=model)
|
||||
else: # Load a custom DiT checkpoint:
|
||||
assert os.path.isfile(model_name), f"Could not find DiT checkpoint at {model_name}"
|
||||
model_ckpt = torch.load(model_name, map_location=lambda storage, loc: storage)
|
||||
model_ckpt = reparameter(model_ckpt, model_name, model=model)
|
||||
return model_ckpt
|
||||
|
||||
|
||||
def download_model(model_name=None, local_path=None, url=None):
|
||||
"""
|
||||
Downloads a pre-trained DiT model from the web.
|
||||
"""
|
||||
if model_name is not None:
|
||||
assert model_name in pretrained_models
|
||||
local_path = f"pretrained_models/{model_name}"
|
||||
web_path = pretrained_models[model_name]
|
||||
else:
|
||||
assert local_path is not None
|
||||
assert url is not None
|
||||
web_path = url
|
||||
if not os.path.isfile(local_path):
|
||||
os.makedirs("pretrained_models", exist_ok=True)
|
||||
dir_name = os.path.dirname(local_path)
|
||||
file_name = os.path.basename(local_path)
|
||||
download_url(web_path, dir_name, file_name)
|
||||
model = torch.load(local_path, map_location=lambda storage, loc: storage)
|
||||
return model
|
||||
|
||||
|
||||
def load_checkpoint(model, ckpt_path, save_as_pt=False, model_name="model", strict=False):
|
||||
if ckpt_path.endswith(".pt") or ckpt_path.endswith(".pth"):
|
||||
state_dict = find_model(ckpt_path, model=model)
|
||||
missing_keys, unexpected_keys = model.load_state_dict(state_dict, strict=strict)
|
||||
get_logger().info("Missing keys: %s", missing_keys)
|
||||
get_logger().info("Unexpected keys: %s", unexpected_keys)
|
||||
elif os.path.isdir(ckpt_path):
|
||||
load_from_sharded_state_dict(model, ckpt_path, model_name, strict=strict)
|
||||
get_logger().info("Model checkpoint loaded from %s", ckpt_path)
|
||||
if save_as_pt:
|
||||
save_path = os.path.join(ckpt_path, model_name + "_ckpt.pt")
|
||||
torch.save(model.state_dict(), save_path)
|
||||
get_logger().info("Model checkpoint saved to %s", save_path)
|
||||
else:
|
||||
raise ValueError(f"Invalid checkpoint path: {ckpt_path}")
|
||||
|
||||
|
||||
def auto_grad_checkpoint(module, *args, **kwargs):
|
||||
if getattr(module, "grad_checkpointing", False):
|
||||
if not isinstance(module, Iterable):
|
||||
return checkpoint(module, *args, use_reentrant=False, **kwargs)
|
||||
gc_step = module[0].grad_checkpointing_step
|
||||
return checkpoint_sequential(module, gc_step, *args, use_reentrant=False, **kwargs)
|
||||
return module(*args, **kwargs)
|
||||
Executable
+769
@@ -0,0 +1,769 @@
|
||||
# Adapted from OpenSora
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# OpenSora: https://github.com/hpcaitech/Open-Sora
|
||||
# --------------------------------------------------------
|
||||
|
||||
import os
|
||||
from typing import Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from diffusers.models import AutoencoderKL, AutoencoderKLTemporalDecoder
|
||||
from einops import rearrange
|
||||
from transformers import PretrainedConfig, PreTrainedModel
|
||||
|
||||
from .utils import load_checkpoint
|
||||
|
||||
|
||||
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, dtype=self.mean.dtype)
|
||||
|
||||
def sample(self):
|
||||
# torch.randn: standard normal distribution
|
||||
x = self.mean + self.std * torch.randn(self.mean.shape).to(device=self.parameters.device, dtype=self.mean.dtype)
|
||||
return x
|
||||
|
||||
def kl(self, other=None):
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.0])
|
||||
else:
|
||||
if other is None: # SCH: assumes other is a standard normal distribution
|
||||
return 0.5 * torch.sum(torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar, dim=[1, 2, 3, 4])
|
||||
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, 4],
|
||||
)
|
||||
|
||||
def nll(self, sample, dims=[1, 2, 3, 4]):
|
||||
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)
|
||||
|
||||
def mode(self):
|
||||
return self.mean
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
def pad_at_dim(t, pad, dim=-1):
|
||||
dims_from_right = (-dim - 1) if dim < 0 else (t.ndim - dim - 1)
|
||||
zeros = (0, 0) * dims_from_right
|
||||
return F.pad(t, (*zeros, *pad), mode="constant")
|
||||
|
||||
|
||||
def exists(v):
|
||||
return v is not None
|
||||
|
||||
|
||||
class CausalConv3d(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
chan_in,
|
||||
chan_out,
|
||||
kernel_size: Union[int, Tuple[int, int, int]],
|
||||
pad_mode="constant",
|
||||
strides=None, # allow custom stride
|
||||
**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 = strides[0] if strides is not None else 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 = strides if strides is not None else (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=self.pad_mode)
|
||||
x = self.conv(x)
|
||||
return x
|
||||
|
||||
|
||||
class ResBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels, # SCH: added
|
||||
filters,
|
||||
conv_fn,
|
||||
activation_fn=nn.SiLU,
|
||||
use_conv_shortcut=False,
|
||||
num_groups=32,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.filters = filters
|
||||
self.activate = activation_fn()
|
||||
self.use_conv_shortcut = use_conv_shortcut
|
||||
|
||||
# SCH: MAGVIT uses GroupNorm by default
|
||||
self.norm1 = nn.GroupNorm(num_groups, in_channels)
|
||||
self.conv1 = conv_fn(in_channels, self.filters, kernel_size=(3, 3, 3), bias=False)
|
||||
self.norm2 = nn.GroupNorm(num_groups, self.filters)
|
||||
self.conv2 = conv_fn(self.filters, self.filters, kernel_size=(3, 3, 3), bias=False)
|
||||
if in_channels != filters:
|
||||
if self.use_conv_shortcut:
|
||||
self.conv3 = conv_fn(in_channels, self.filters, kernel_size=(3, 3, 3), bias=False)
|
||||
else:
|
||||
self.conv3 = conv_fn(in_channels, self.filters, kernel_size=(1, 1, 1), bias=False)
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
x = self.norm1(x)
|
||||
x = self.activate(x)
|
||||
x = self.conv1(x)
|
||||
x = self.norm2(x)
|
||||
x = self.activate(x)
|
||||
x = self.conv2(x)
|
||||
if self.in_channels != self.filters: # SCH: ResBlock X->Y
|
||||
residual = self.conv3(residual)
|
||||
return x + residual
|
||||
|
||||
|
||||
def get_activation_fn(activation):
|
||||
if activation == "relu":
|
||||
activation_fn = nn.ReLU
|
||||
elif activation == "swish":
|
||||
activation_fn = nn.SiLU
|
||||
else:
|
||||
raise NotImplementedError
|
||||
return activation_fn
|
||||
|
||||
|
||||
class Encoder(nn.Module):
|
||||
"""Encoder Blocks."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_out_channels=4,
|
||||
latent_embed_dim=512, # num channels for latent vector
|
||||
filters=128,
|
||||
num_res_blocks=4,
|
||||
channel_multipliers=(1, 2, 2, 4),
|
||||
temporal_downsample=(False, True, True),
|
||||
num_groups=32, # for nn.GroupNorm
|
||||
activation_fn="swish",
|
||||
):
|
||||
super().__init__()
|
||||
self.filters = filters
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.num_blocks = len(channel_multipliers)
|
||||
self.channel_multipliers = channel_multipliers
|
||||
self.temporal_downsample = temporal_downsample
|
||||
self.num_groups = num_groups
|
||||
self.embedding_dim = latent_embed_dim
|
||||
|
||||
self.activation_fn = get_activation_fn(activation_fn)
|
||||
self.activate = self.activation_fn()
|
||||
self.conv_fn = CausalConv3d
|
||||
self.block_args = dict(
|
||||
conv_fn=self.conv_fn,
|
||||
activation_fn=self.activation_fn,
|
||||
use_conv_shortcut=False,
|
||||
num_groups=self.num_groups,
|
||||
)
|
||||
|
||||
# first layer conv
|
||||
self.conv_in = self.conv_fn(
|
||||
in_out_channels,
|
||||
filters,
|
||||
kernel_size=(3, 3, 3),
|
||||
bias=False,
|
||||
)
|
||||
|
||||
# ResBlocks and conv downsample
|
||||
self.block_res_blocks = nn.ModuleList([])
|
||||
self.conv_blocks = nn.ModuleList([])
|
||||
|
||||
filters = self.filters
|
||||
prev_filters = filters # record for in_channels
|
||||
for i in range(self.num_blocks):
|
||||
filters = self.filters * self.channel_multipliers[i]
|
||||
block_items = nn.ModuleList([])
|
||||
for _ in range(self.num_res_blocks):
|
||||
block_items.append(ResBlock(prev_filters, filters, **self.block_args))
|
||||
prev_filters = filters # update in_channels
|
||||
self.block_res_blocks.append(block_items)
|
||||
|
||||
if i < self.num_blocks - 1:
|
||||
if self.temporal_downsample[i]:
|
||||
t_stride = 2 if self.temporal_downsample[i] else 1
|
||||
s_stride = 1
|
||||
self.conv_blocks.append(
|
||||
self.conv_fn(
|
||||
prev_filters, filters, kernel_size=(3, 3, 3), strides=(t_stride, s_stride, s_stride)
|
||||
)
|
||||
)
|
||||
prev_filters = filters # update in_channels
|
||||
else:
|
||||
# if no t downsample, don't add since this does nothing for pipeline models
|
||||
self.conv_blocks.append(nn.Identity(prev_filters)) # Identity
|
||||
prev_filters = filters # update in_channels
|
||||
|
||||
# last layer res block
|
||||
self.res_blocks = nn.ModuleList([])
|
||||
for _ in range(self.num_res_blocks):
|
||||
self.res_blocks.append(ResBlock(prev_filters, filters, **self.block_args))
|
||||
prev_filters = filters # update in_channels
|
||||
|
||||
# MAGVIT uses Group Normalization
|
||||
self.norm1 = nn.GroupNorm(self.num_groups, prev_filters)
|
||||
|
||||
self.conv2 = self.conv_fn(prev_filters, self.embedding_dim, kernel_size=(1, 1, 1), padding="same")
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv_in(x)
|
||||
|
||||
for i in range(self.num_blocks):
|
||||
for j in range(self.num_res_blocks):
|
||||
x = self.block_res_blocks[i][j](x)
|
||||
if i < self.num_blocks - 1:
|
||||
x = self.conv_blocks[i](x)
|
||||
for i in range(self.num_res_blocks):
|
||||
x = self.res_blocks[i](x)
|
||||
|
||||
x = self.norm1(x)
|
||||
x = self.activate(x)
|
||||
x = self.conv2(x)
|
||||
return x
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
"""Decoder Blocks."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_out_channels=4,
|
||||
latent_embed_dim=512,
|
||||
filters=128,
|
||||
num_res_blocks=4,
|
||||
channel_multipliers=(1, 2, 2, 4),
|
||||
temporal_downsample=(False, True, True),
|
||||
num_groups=32, # for nn.GroupNorm
|
||||
activation_fn="swish",
|
||||
):
|
||||
super().__init__()
|
||||
self.filters = filters
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.num_blocks = len(channel_multipliers)
|
||||
self.channel_multipliers = channel_multipliers
|
||||
self.temporal_downsample = temporal_downsample
|
||||
self.num_groups = num_groups
|
||||
self.embedding_dim = latent_embed_dim
|
||||
self.s_stride = 1
|
||||
|
||||
self.activation_fn = get_activation_fn(activation_fn)
|
||||
self.activate = self.activation_fn()
|
||||
self.conv_fn = CausalConv3d
|
||||
self.block_args = dict(
|
||||
conv_fn=self.conv_fn,
|
||||
activation_fn=self.activation_fn,
|
||||
use_conv_shortcut=False,
|
||||
num_groups=self.num_groups,
|
||||
)
|
||||
|
||||
filters = self.filters * self.channel_multipliers[-1]
|
||||
prev_filters = filters
|
||||
|
||||
# last conv
|
||||
self.conv1 = self.conv_fn(self.embedding_dim, filters, kernel_size=(3, 3, 3), bias=True)
|
||||
|
||||
# last layer res block
|
||||
self.res_blocks = nn.ModuleList([])
|
||||
for _ in range(self.num_res_blocks):
|
||||
self.res_blocks.append(ResBlock(filters, filters, **self.block_args))
|
||||
|
||||
# ResBlocks and conv upsample
|
||||
self.block_res_blocks = nn.ModuleList([])
|
||||
self.num_blocks = len(self.channel_multipliers)
|
||||
self.conv_blocks = nn.ModuleList([])
|
||||
# reverse to keep track of the in_channels, but append also in a reverse direction
|
||||
for i in reversed(range(self.num_blocks)):
|
||||
filters = self.filters * self.channel_multipliers[i]
|
||||
# resblock handling
|
||||
block_items = nn.ModuleList([])
|
||||
for _ in range(self.num_res_blocks):
|
||||
block_items.append(ResBlock(prev_filters, filters, **self.block_args))
|
||||
prev_filters = filters # SCH: update in_channels
|
||||
self.block_res_blocks.insert(0, block_items) # SCH: append in front
|
||||
|
||||
# conv blocks with upsampling
|
||||
if i > 0:
|
||||
if self.temporal_downsample[i - 1]:
|
||||
t_stride = 2 if self.temporal_downsample[i - 1] else 1
|
||||
# SCH: T-Causal Conv 3x3x3, f -> (t_stride * 2 * 2) * f, depth to space t_stride x 2 x 2
|
||||
self.conv_blocks.insert(
|
||||
0,
|
||||
self.conv_fn(
|
||||
prev_filters, prev_filters * t_stride * self.s_stride * self.s_stride, kernel_size=(3, 3, 3)
|
||||
),
|
||||
)
|
||||
else:
|
||||
self.conv_blocks.insert(
|
||||
0,
|
||||
nn.Identity(prev_filters),
|
||||
)
|
||||
|
||||
self.norm1 = nn.GroupNorm(self.num_groups, prev_filters)
|
||||
|
||||
self.conv_out = self.conv_fn(filters, in_out_channels, 3)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv1(x)
|
||||
for i in range(self.num_res_blocks):
|
||||
x = self.res_blocks[i](x)
|
||||
for i in reversed(range(self.num_blocks)):
|
||||
for j in range(self.num_res_blocks):
|
||||
x = self.block_res_blocks[i][j](x)
|
||||
if i > 0:
|
||||
t_stride = 2 if self.temporal_downsample[i - 1] else 1
|
||||
x = self.conv_blocks[i - 1](x)
|
||||
x = rearrange(
|
||||
x,
|
||||
"B (C ts hs ws) T H W -> B C (T ts) (H hs) (W ws)",
|
||||
ts=t_stride,
|
||||
hs=self.s_stride,
|
||||
ws=self.s_stride,
|
||||
)
|
||||
|
||||
x = self.norm1(x)
|
||||
x = self.activate(x)
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class VAE_Temporal(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_out_channels=4,
|
||||
latent_embed_dim=4,
|
||||
embed_dim=4,
|
||||
filters=128,
|
||||
num_res_blocks=4,
|
||||
channel_multipliers=(1, 2, 2, 4),
|
||||
temporal_downsample=(True, True, False),
|
||||
num_groups=32, # for nn.GroupNorm
|
||||
activation_fn="swish",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.time_downsample_factor = 2 ** sum(temporal_downsample)
|
||||
# self.time_padding = self.time_downsample_factor - 1
|
||||
self.patch_size = (self.time_downsample_factor, 1, 1)
|
||||
self.out_channels = in_out_channels
|
||||
|
||||
# NOTE: following MAGVIT, conv in bias=False in encoder first conv
|
||||
self.encoder = Encoder(
|
||||
in_out_channels=in_out_channels,
|
||||
latent_embed_dim=latent_embed_dim * 2,
|
||||
filters=filters,
|
||||
num_res_blocks=num_res_blocks,
|
||||
channel_multipliers=channel_multipliers,
|
||||
temporal_downsample=temporal_downsample,
|
||||
num_groups=num_groups, # for nn.GroupNorm
|
||||
activation_fn=activation_fn,
|
||||
)
|
||||
self.quant_conv = CausalConv3d(2 * latent_embed_dim, 2 * embed_dim, 1)
|
||||
|
||||
self.post_quant_conv = CausalConv3d(embed_dim, latent_embed_dim, 1)
|
||||
self.decoder = Decoder(
|
||||
in_out_channels=in_out_channels,
|
||||
latent_embed_dim=latent_embed_dim,
|
||||
filters=filters,
|
||||
num_res_blocks=num_res_blocks,
|
||||
channel_multipliers=channel_multipliers,
|
||||
temporal_downsample=temporal_downsample,
|
||||
num_groups=num_groups, # for nn.GroupNorm
|
||||
activation_fn=activation_fn,
|
||||
)
|
||||
|
||||
def get_latent_size(self, input_size):
|
||||
latent_size = []
|
||||
for i in range(3):
|
||||
if input_size[i] is None:
|
||||
lsize = None
|
||||
elif i == 0:
|
||||
time_padding = (
|
||||
0
|
||||
if (input_size[i] % self.time_downsample_factor == 0)
|
||||
else self.time_downsample_factor - input_size[i] % self.time_downsample_factor
|
||||
)
|
||||
lsize = (input_size[i] + time_padding) // self.patch_size[i]
|
||||
else:
|
||||
lsize = input_size[i] // self.patch_size[i]
|
||||
latent_size.append(lsize)
|
||||
return latent_size
|
||||
|
||||
def encode(self, x):
|
||||
time_padding = (
|
||||
0
|
||||
if (x.shape[2] % self.time_downsample_factor == 0)
|
||||
else self.time_downsample_factor - x.shape[2] % self.time_downsample_factor
|
||||
)
|
||||
x = pad_at_dim(x, (time_padding, 0), dim=2)
|
||||
encoded_feature = self.encoder(x)
|
||||
moments = self.quant_conv(encoded_feature).to(x.dtype)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
return posterior
|
||||
|
||||
def decode(self, z, num_frames=None):
|
||||
time_padding = (
|
||||
0
|
||||
if (num_frames % self.time_downsample_factor == 0)
|
||||
else self.time_downsample_factor - num_frames % self.time_downsample_factor
|
||||
)
|
||||
z = self.post_quant_conv(z)
|
||||
x = self.decoder(z)
|
||||
x = x[:, :, time_padding:]
|
||||
return x
|
||||
|
||||
def forward(self, x, sample_posterior=True):
|
||||
posterior = self.encode(x)
|
||||
if sample_posterior:
|
||||
z = posterior.sample()
|
||||
else:
|
||||
z = posterior.mode()
|
||||
recon_video = self.decode(z, num_frames=x.shape[2])
|
||||
return recon_video, posterior, z
|
||||
|
||||
|
||||
def VAE_Temporal_SD(from_pretrained=None, **kwargs):
|
||||
model = VAE_Temporal(
|
||||
in_out_channels=4,
|
||||
latent_embed_dim=4,
|
||||
embed_dim=4,
|
||||
filters=128,
|
||||
num_res_blocks=4,
|
||||
channel_multipliers=(1, 2, 2, 4),
|
||||
temporal_downsample=(False, True, True),
|
||||
**kwargs,
|
||||
)
|
||||
if from_pretrained is not None:
|
||||
load_checkpoint(model, from_pretrained)
|
||||
return model
|
||||
|
||||
|
||||
class VideoAutoencoderKL(nn.Module):
|
||||
def __init__(
|
||||
self, from_pretrained=None, micro_batch_size=None, cache_dir=None, local_files_only=False, subfolder=None
|
||||
):
|
||||
super().__init__()
|
||||
self.module = AutoencoderKL.from_pretrained(
|
||||
from_pretrained,
|
||||
cache_dir=cache_dir,
|
||||
local_files_only=local_files_only,
|
||||
subfolder=subfolder,
|
||||
)
|
||||
self.out_channels = self.module.config.latent_channels
|
||||
self.patch_size = (1, 8, 8)
|
||||
self.micro_batch_size = micro_batch_size
|
||||
|
||||
def encode(self, x):
|
||||
# x: (B, C, T, H, W)
|
||||
B = x.shape[0]
|
||||
x = rearrange(x, "B C T H W -> (B T) C H W")
|
||||
|
||||
if self.micro_batch_size is None:
|
||||
x = self.module.encode(x).latent_dist.sample().mul_(0.18215)
|
||||
else:
|
||||
# NOTE: cannot be used for training
|
||||
bs = self.micro_batch_size
|
||||
x_out = []
|
||||
for i in range(0, x.shape[0], bs):
|
||||
x_bs = x[i : i + bs]
|
||||
x_bs = self.module.encode(x_bs).latent_dist.sample().mul_(0.18215)
|
||||
x_out.append(x_bs)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
x = rearrange(x, "(B T) C H W -> B C T H W", B=B)
|
||||
return x
|
||||
|
||||
def decode(self, x, **kwargs):
|
||||
# x: (B, C, T, H, W)
|
||||
B = x.shape[0]
|
||||
x = rearrange(x, "B C T H W -> (B T) C H W")
|
||||
if self.micro_batch_size is None:
|
||||
x = self.module.decode(x / 0.18215).sample
|
||||
else:
|
||||
# NOTE: cannot be used for training
|
||||
bs = self.micro_batch_size
|
||||
x_out = []
|
||||
for i in range(0, x.shape[0], bs):
|
||||
x_bs = x[i : i + bs]
|
||||
x_bs = self.module.decode(x_bs / 0.18215).sample
|
||||
x_out.append(x_bs)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
x = rearrange(x, "(B T) C H W -> B C T H W", B=B)
|
||||
return x
|
||||
|
||||
def get_latent_size(self, input_size):
|
||||
latent_size = []
|
||||
for i in range(3):
|
||||
# assert (
|
||||
# input_size[i] is None or input_size[i] % self.patch_size[i] == 0
|
||||
# ), "Input size must be divisible by patch size"
|
||||
latent_size.append(input_size[i] // self.patch_size[i] if input_size[i] is not None else None)
|
||||
return latent_size
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
|
||||
class VideoAutoencoderKLTemporalDecoder(nn.Module):
|
||||
def __init__(self, from_pretrained=None, cache_dir=None, local_files_only=False):
|
||||
super().__init__()
|
||||
self.module = AutoencoderKLTemporalDecoder.from_pretrained(
|
||||
from_pretrained, cache_dir=cache_dir, local_files_only=local_files_only
|
||||
)
|
||||
self.out_channels = self.module.config.latent_channels
|
||||
self.patch_size = (1, 8, 8)
|
||||
|
||||
def encode(self, x):
|
||||
raise NotImplementedError
|
||||
|
||||
def decode(self, x, **kwargs):
|
||||
B, _, T = x.shape[:3]
|
||||
x = rearrange(x, "B C T H W -> (B T) C H W")
|
||||
x = self.module.decode(x / 0.18215, num_frames=T).sample
|
||||
x = rearrange(x, "(B T) C H W -> B C T H W", B=B)
|
||||
return x
|
||||
|
||||
def get_latent_size(self, input_size):
|
||||
latent_size = []
|
||||
for i in range(3):
|
||||
# assert (
|
||||
# input_size[i] is None or input_size[i] % self.patch_size[i] == 0
|
||||
# ), "Input size must be divisible by patch size"
|
||||
latent_size.append(input_size[i] // self.patch_size[i] if input_size[i] is not None else None)
|
||||
return latent_size
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
|
||||
class VideoAutoencoderPipelineConfig(PretrainedConfig):
|
||||
model_type = "VideoAutoencoderPipeline"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae_2d=None,
|
||||
vae_temporal=None,
|
||||
from_pretrained=None,
|
||||
freeze_vae_2d=False,
|
||||
cal_loss=False,
|
||||
micro_frame_size=None,
|
||||
shift=0.0,
|
||||
scale=1.0,
|
||||
**kwargs,
|
||||
):
|
||||
self.vae_2d = vae_2d
|
||||
self.vae_temporal = vae_temporal
|
||||
self.from_pretrained = from_pretrained
|
||||
self.freeze_vae_2d = freeze_vae_2d
|
||||
self.cal_loss = cal_loss
|
||||
self.micro_frame_size = micro_frame_size
|
||||
self.shift = shift
|
||||
self.scale = scale
|
||||
super().__init__(**kwargs)
|
||||
|
||||
|
||||
class VideoAutoencoderPipeline(PreTrainedModel):
|
||||
config_class = VideoAutoencoderPipelineConfig
|
||||
|
||||
def __init__(self, config: VideoAutoencoderPipelineConfig):
|
||||
super().__init__(config=config)
|
||||
self.spatial_vae = VideoAutoencoderKL(
|
||||
from_pretrained="PixArt-alpha/pixart_sigma_sdxlvae_T5_diffusers",
|
||||
local_files_only=False,
|
||||
micro_batch_size=4,
|
||||
subfolder="vae",
|
||||
)
|
||||
self.temporal_vae = VAE_Temporal_SD(from_pretrained=None)
|
||||
self.cal_loss = config.cal_loss
|
||||
self.micro_frame_size = config.micro_frame_size
|
||||
self.micro_z_frame_size = self.temporal_vae.get_latent_size([config.micro_frame_size, None, None])[0]
|
||||
|
||||
if config.freeze_vae_2d:
|
||||
for param in self.spatial_vae.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
self.out_channels = self.temporal_vae.out_channels
|
||||
|
||||
# normalization parameters
|
||||
scale = torch.tensor(config.scale)
|
||||
shift = torch.tensor(config.shift)
|
||||
if len(scale.shape) > 0:
|
||||
scale = scale[None, :, None, None, None]
|
||||
if len(shift.shape) > 0:
|
||||
shift = shift[None, :, None, None, None]
|
||||
self.register_buffer("scale", scale)
|
||||
self.register_buffer("shift", shift)
|
||||
|
||||
def encode(self, x):
|
||||
x_z = self.spatial_vae.encode(x)
|
||||
|
||||
if self.micro_frame_size is None:
|
||||
posterior = self.temporal_vae.encode(x_z)
|
||||
z = posterior.sample()
|
||||
else:
|
||||
z_list = []
|
||||
for i in range(0, x_z.shape[2], self.micro_frame_size):
|
||||
x_z_bs = x_z[:, :, i : i + self.micro_frame_size]
|
||||
posterior = self.temporal_vae.encode(x_z_bs)
|
||||
z_list.append(posterior.sample())
|
||||
z = torch.cat(z_list, dim=2)
|
||||
|
||||
if self.cal_loss:
|
||||
return z, posterior, x_z
|
||||
else:
|
||||
return (z - self.shift) / self.scale
|
||||
|
||||
def decode(self, z, num_frames=None):
|
||||
if not self.cal_loss:
|
||||
z = z * self.scale.to(z.dtype) + self.shift.to(z.dtype)
|
||||
|
||||
if self.micro_frame_size is None:
|
||||
x_z = self.temporal_vae.decode(z, num_frames=num_frames)
|
||||
x = self.spatial_vae.decode(x_z)
|
||||
else:
|
||||
x_z_list = []
|
||||
for i in range(0, z.size(2), self.micro_z_frame_size):
|
||||
z_bs = z[:, :, i : i + self.micro_z_frame_size]
|
||||
x_z_bs = self.temporal_vae.decode(z_bs, num_frames=min(self.micro_frame_size, num_frames))
|
||||
x_z_list.append(x_z_bs)
|
||||
num_frames -= self.micro_frame_size
|
||||
x_z = torch.cat(x_z_list, dim=2)
|
||||
x = self.spatial_vae.decode(x_z)
|
||||
|
||||
if self.cal_loss:
|
||||
return x, x_z
|
||||
else:
|
||||
return x
|
||||
|
||||
def forward(self, x):
|
||||
assert self.cal_loss, "This method is only available when cal_loss is True"
|
||||
z, posterior, x_z = self.encode(x)
|
||||
x_rec, x_z_rec = self.decode(z, num_frames=x_z.shape[2])
|
||||
return x_rec, x_z_rec, z, posterior, x_z
|
||||
|
||||
def get_latent_size(self, input_size):
|
||||
if self.micro_frame_size is None or input_size[0] is None:
|
||||
return self.temporal_vae.get_latent_size(self.spatial_vae.get_latent_size(input_size))
|
||||
else:
|
||||
sub_input_size = [self.micro_frame_size, input_size[1], input_size[2]]
|
||||
sub_latent_size = self.temporal_vae.get_latent_size(self.spatial_vae.get_latent_size(sub_input_size))
|
||||
sub_latent_size[0] = sub_latent_size[0] * (input_size[0] // self.micro_frame_size)
|
||||
remain_temporal_size = [input_size[0] % self.micro_frame_size, None, None]
|
||||
if remain_temporal_size[0] > 0:
|
||||
remain_size = self.temporal_vae.get_latent_size(remain_temporal_size)
|
||||
sub_latent_size[0] += remain_size[0]
|
||||
return sub_latent_size
|
||||
|
||||
def get_temporal_last_layer(self):
|
||||
return self.temporal_vae.decoder.conv_out.conv.weight
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
|
||||
def OpenSoraVAE_V1_2(
|
||||
micro_batch_size=4,
|
||||
micro_frame_size=17,
|
||||
from_pretrained=None,
|
||||
local_files_only=False,
|
||||
freeze_vae_2d=False,
|
||||
cal_loss=False,
|
||||
):
|
||||
vae_2d = dict(
|
||||
type="VideoAutoencoderKL",
|
||||
from_pretrained="PixArt-alpha/pixart_sigma_sdxlvae_T5_diffusers",
|
||||
subfolder="vae",
|
||||
micro_batch_size=micro_batch_size,
|
||||
local_files_only=local_files_only,
|
||||
)
|
||||
vae_temporal = dict(
|
||||
type="VAE_Temporal_SD",
|
||||
from_pretrained=None,
|
||||
)
|
||||
shift = (-0.10, 0.34, 0.27, 0.98)
|
||||
scale = (3.85, 2.32, 2.33, 3.06)
|
||||
kwargs = dict(
|
||||
vae_2d=vae_2d,
|
||||
vae_temporal=vae_temporal,
|
||||
freeze_vae_2d=freeze_vae_2d,
|
||||
cal_loss=cal_loss,
|
||||
micro_frame_size=micro_frame_size,
|
||||
shift=shift,
|
||||
scale=scale,
|
||||
)
|
||||
|
||||
if from_pretrained is not None and not os.path.isdir(from_pretrained):
|
||||
model = VideoAutoencoderPipeline.from_pretrained(from_pretrained, **kwargs)
|
||||
else:
|
||||
config = VideoAutoencoderPipelineConfig(**kwargs)
|
||||
model = VideoAutoencoderPipeline(config)
|
||||
|
||||
if from_pretrained:
|
||||
load_checkpoint(model, from_pretrained)
|
||||
return model
|
||||
Executable
+5
@@ -0,0 +1,5 @@
|
||||
from .ae import ae_stride_config, getae_wrapper
|
||||
from .latte import LatteT2V
|
||||
from .pipeline import VideoGenPipeline
|
||||
|
||||
__all__ = ["VideoGenPipeline", "ae_stride_config", "getae_wrapper", "LatteT2V"]
|
||||
Executable
+857
@@ -0,0 +1,857 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
import glob
|
||||
import importlib
|
||||
import os
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import ConfigMixin, ModelMixin
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from einops import rearrange
|
||||
from torch import nn
|
||||
|
||||
|
||||
def Normalize(in_channels, num_groups=32):
|
||||
return torch.nn.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||
|
||||
|
||||
def tensor_to_video(x):
|
||||
x = x.detach().cpu()
|
||||
x = torch.clamp(x, -1, 1)
|
||||
x = (x + 1) / 2
|
||||
x = x.permute(1, 0, 2, 3).float().numpy() # c t h w ->
|
||||
x = (255 * x).astype(np.uint8)
|
||||
return x
|
||||
|
||||
|
||||
def nonlinearity(x):
|
||||
return x * torch.sigmoid(x)
|
||||
|
||||
|
||||
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.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.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 resolve_str_to_obj(str_val, append=True):
|
||||
if append:
|
||||
str_val = "opendit.models.opensora_plan.modules." + str_val
|
||||
if "opensora.models.ae.videobase." in str_val:
|
||||
str_val = str_val.replace("opensora.models.ae.videobase.", "opendit.models.opensora_plan.")
|
||||
module_name, class_name = str_val.rsplit(".", 1)
|
||||
module = importlib.import_module(module_name)
|
||||
return getattr(module, class_name)
|
||||
|
||||
|
||||
class VideoBaseAE_PL(ModelMixin, ConfigMixin):
|
||||
config_name = "config.json"
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def encode(self, x: torch.Tensor, *args, **kwargs):
|
||||
pass
|
||||
|
||||
def decode(self, encoding: torch.Tensor, *args, **kwargs):
|
||||
pass
|
||||
|
||||
@property
|
||||
def num_training_steps(self) -> int:
|
||||
"""Total training steps inferred from datamodule and devices."""
|
||||
if self.trainer.max_steps:
|
||||
return self.trainer.max_steps
|
||||
|
||||
limit_batches = self.trainer.limit_train_batches
|
||||
batches = len(self.train_dataloader())
|
||||
batches = min(batches, limit_batches) if isinstance(limit_batches, int) else int(limit_batches * batches)
|
||||
|
||||
num_devices = max(1, self.trainer.num_gpus, self.trainer.num_processes)
|
||||
if self.trainer.tpu_cores:
|
||||
num_devices = max(num_devices, self.trainer.tpu_cores)
|
||||
|
||||
effective_accum = self.trainer.accumulate_grad_batches * num_devices
|
||||
return (batches // effective_accum) * self.trainer.max_epochs
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, pretrained_model_name_or_path: Optional[Union[str, os.PathLike]], **kwargs):
|
||||
ckpt_files = glob.glob(os.path.join(pretrained_model_name_or_path, "*.ckpt"))
|
||||
if ckpt_files:
|
||||
# Adapt to PyTorch Lightning
|
||||
last_ckpt_file = ckpt_files[-1]
|
||||
config_file = os.path.join(pretrained_model_name_or_path, cls.config_name)
|
||||
model = cls.from_config(config_file)
|
||||
print("init from {}".format(last_ckpt_file))
|
||||
model.init_from_ckpt(last_ckpt_file)
|
||||
return model
|
||||
else:
|
||||
print(f"Loading model from {pretrained_model_name_or_path}")
|
||||
return super().from_pretrained(pretrained_model_name_or_path, **kwargs)
|
||||
|
||||
|
||||
class Encoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
z_channels: int,
|
||||
hidden_size: int,
|
||||
hidden_size_mult: Tuple[int] = (1, 2, 4, 4),
|
||||
attn_resolutions: Tuple[int] = (16,),
|
||||
conv_in: str = "Conv2d",
|
||||
conv_out: str = "CasualConv3d",
|
||||
attention: str = "AttnBlock",
|
||||
resnet_blocks: Tuple[str] = (
|
||||
"ResnetBlock2D",
|
||||
"ResnetBlock2D",
|
||||
"ResnetBlock2D",
|
||||
"ResnetBlock3D",
|
||||
),
|
||||
spatial_downsample: Tuple[str] = (
|
||||
"Downsample",
|
||||
"Downsample",
|
||||
"Downsample",
|
||||
"",
|
||||
),
|
||||
temporal_downsample: Tuple[str] = ("", "", "TimeDownsampleRes2x", ""),
|
||||
mid_resnet: str = "ResnetBlock3D",
|
||||
dropout: float = 0.0,
|
||||
resolution: int = 256,
|
||||
num_res_blocks: int = 2,
|
||||
double_z: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
assert len(resnet_blocks) == len(hidden_size_mult), print(hidden_size_mult, resnet_blocks)
|
||||
# ---- Config ----
|
||||
self.num_resolutions = len(hidden_size_mult)
|
||||
self.resolution = resolution
|
||||
self.num_res_blocks = num_res_blocks
|
||||
|
||||
# ---- In ----
|
||||
self.conv_in = resolve_str_to_obj(conv_in)(3, hidden_size, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
# ---- Downsample ----
|
||||
curr_res = resolution
|
||||
in_ch_mult = (1,) + tuple(hidden_size_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 = hidden_size * in_ch_mult[i_level]
|
||||
block_out = hidden_size * hidden_size_mult[i_level]
|
||||
for i_block in range(self.num_res_blocks):
|
||||
block.append(
|
||||
resolve_str_to_obj(resnet_blocks[i_level])(
|
||||
in_channels=block_in,
|
||||
out_channels=block_out,
|
||||
dropout=dropout,
|
||||
)
|
||||
)
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(resolve_str_to_obj(attention)(block_in))
|
||||
down = nn.Module()
|
||||
down.block = block
|
||||
down.attn = attn
|
||||
if spatial_downsample[i_level]:
|
||||
down.downsample = resolve_str_to_obj(spatial_downsample[i_level])(block_in, block_in)
|
||||
curr_res = curr_res // 2
|
||||
if temporal_downsample[i_level]:
|
||||
down.time_downsample = resolve_str_to_obj(temporal_downsample[i_level])(block_in, block_in)
|
||||
self.down.append(down)
|
||||
|
||||
# ---- Mid ----
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = resolve_str_to_obj(mid_resnet)(
|
||||
in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
dropout=dropout,
|
||||
)
|
||||
self.mid.attn_1 = resolve_str_to_obj(attention)(block_in)
|
||||
self.mid.block_2 = resolve_str_to_obj(mid_resnet)(
|
||||
in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
dropout=dropout,
|
||||
)
|
||||
# ---- Out ----
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = resolve_str_to_obj(conv_out)(
|
||||
block_in,
|
||||
2 * z_channels if double_z else z_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
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])
|
||||
if len(self.down[i_level].attn) > 0:
|
||||
h = self.down[i_level].attn[i_block](h)
|
||||
hs.append(h)
|
||||
if hasattr(self.down[i_level], "downsample"):
|
||||
hs.append(self.down[i_level].downsample(hs[-1]))
|
||||
if hasattr(self.down[i_level], "time_downsample"):
|
||||
hs_down = self.down[i_level].time_downsample(hs[-1])
|
||||
hs.append(hs_down)
|
||||
|
||||
h = self.mid.block_1(h)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h)
|
||||
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
return h
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
z_channels: int,
|
||||
hidden_size: int,
|
||||
hidden_size_mult: Tuple[int] = (1, 2, 4, 4),
|
||||
attn_resolutions: Tuple[int] = (16,),
|
||||
conv_in: str = "Conv2d",
|
||||
conv_out: str = "CasualConv3d",
|
||||
attention: str = "AttnBlock",
|
||||
resnet_blocks: Tuple[str] = (
|
||||
"ResnetBlock3D",
|
||||
"ResnetBlock3D",
|
||||
"ResnetBlock3D",
|
||||
"ResnetBlock3D",
|
||||
),
|
||||
spatial_upsample: Tuple[str] = (
|
||||
"",
|
||||
"SpatialUpsample2x",
|
||||
"SpatialUpsample2x",
|
||||
"SpatialUpsample2x",
|
||||
),
|
||||
temporal_upsample: Tuple[str] = ("", "", "", "TimeUpsampleRes2x"),
|
||||
mid_resnet: str = "ResnetBlock3D",
|
||||
dropout: float = 0.0,
|
||||
resolution: int = 256,
|
||||
num_res_blocks: int = 2,
|
||||
):
|
||||
super().__init__()
|
||||
# ---- Config ----
|
||||
self.num_resolutions = len(hidden_size_mult)
|
||||
self.resolution = resolution
|
||||
self.num_res_blocks = num_res_blocks
|
||||
|
||||
# ---- In ----
|
||||
block_in = hidden_size * hidden_size_mult[self.num_resolutions - 1]
|
||||
curr_res = resolution // 2 ** (self.num_resolutions - 1)
|
||||
self.conv_in = resolve_str_to_obj(conv_in)(z_channels, block_in, kernel_size=3, padding=1)
|
||||
|
||||
# ---- Mid ----
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = resolve_str_to_obj(mid_resnet)(
|
||||
in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
dropout=dropout,
|
||||
)
|
||||
self.mid.attn_1 = resolve_str_to_obj(attention)(block_in)
|
||||
self.mid.block_2 = resolve_str_to_obj(mid_resnet)(
|
||||
in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
dropout=dropout,
|
||||
)
|
||||
|
||||
# ---- Upsample ----
|
||||
self.up = nn.ModuleList()
|
||||
for i_level in reversed(range(self.num_resolutions)):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_out = hidden_size * hidden_size_mult[i_level]
|
||||
for i_block in range(self.num_res_blocks + 1):
|
||||
block.append(
|
||||
resolve_str_to_obj(resnet_blocks[i_level])(
|
||||
in_channels=block_in,
|
||||
out_channels=block_out,
|
||||
dropout=dropout,
|
||||
)
|
||||
)
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(resolve_str_to_obj(attention)(block_in))
|
||||
up = nn.Module()
|
||||
up.block = block
|
||||
up.attn = attn
|
||||
if spatial_upsample[i_level]:
|
||||
up.upsample = resolve_str_to_obj(spatial_upsample[i_level])(block_in, block_in)
|
||||
curr_res = curr_res * 2
|
||||
if temporal_upsample[i_level]:
|
||||
up.time_upsample = resolve_str_to_obj(temporal_upsample[i_level])(block_in, block_in)
|
||||
self.up.insert(0, up)
|
||||
|
||||
# ---- Out ----
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = resolve_str_to_obj(conv_out)(block_in, 3, kernel_size=3, padding=1)
|
||||
|
||||
def forward(self, z):
|
||||
h = self.conv_in(z)
|
||||
h = self.mid.block_1(h)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h)
|
||||
|
||||
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)
|
||||
if len(self.up[i_level].attn) > 0:
|
||||
h = self.up[i_level].attn[i_block](h)
|
||||
if hasattr(self.up[i_level], "upsample"):
|
||||
h = self.up[i_level].upsample(h)
|
||||
if hasattr(self.up[i_level], "time_upsample"):
|
||||
h = self.up[i_level].time_upsample(h)
|
||||
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
return h
|
||||
|
||||
|
||||
class CausalVAEModel(VideoBaseAE_PL):
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
lr: float = 1e-5,
|
||||
hidden_size: int = 128,
|
||||
z_channels: int = 4,
|
||||
hidden_size_mult: Tuple[int] = (1, 2, 4, 4),
|
||||
attn_resolutions: Tuple[int] = [],
|
||||
dropout: float = 0.0,
|
||||
resolution: int = 256,
|
||||
double_z: bool = True,
|
||||
embed_dim: int = 4,
|
||||
num_res_blocks: int = 2,
|
||||
loss_type: str = "opensora.models.ae.videobase.losses.LPIPSWithDiscriminator",
|
||||
loss_params: dict = {
|
||||
"kl_weight": 0.000001,
|
||||
"logvar_init": 0.0,
|
||||
"disc_start": 2001,
|
||||
"disc_weight": 0.5,
|
||||
},
|
||||
q_conv: str = "CausalConv3d",
|
||||
encoder_conv_in: str = "CausalConv3d",
|
||||
encoder_conv_out: str = "CausalConv3d",
|
||||
encoder_attention: str = "AttnBlock3D",
|
||||
encoder_resnet_blocks: Tuple[str] = (
|
||||
"ResnetBlock3D",
|
||||
"ResnetBlock3D",
|
||||
"ResnetBlock3D",
|
||||
"ResnetBlock3D",
|
||||
),
|
||||
encoder_spatial_downsample: Tuple[str] = (
|
||||
"SpatialDownsample2x",
|
||||
"SpatialDownsample2x",
|
||||
"SpatialDownsample2x",
|
||||
"",
|
||||
),
|
||||
encoder_temporal_downsample: Tuple[str] = (
|
||||
"",
|
||||
"TimeDownsample2x",
|
||||
"TimeDownsample2x",
|
||||
"",
|
||||
),
|
||||
encoder_mid_resnet: str = "ResnetBlock3D",
|
||||
decoder_conv_in: str = "CausalConv3d",
|
||||
decoder_conv_out: str = "CausalConv3d",
|
||||
decoder_attention: str = "AttnBlock3D",
|
||||
decoder_resnet_blocks: Tuple[str] = (
|
||||
"ResnetBlock3D",
|
||||
"ResnetBlock3D",
|
||||
"ResnetBlock3D",
|
||||
"ResnetBlock3D",
|
||||
),
|
||||
decoder_spatial_upsample: Tuple[str] = (
|
||||
"",
|
||||
"SpatialUpsample2x",
|
||||
"SpatialUpsample2x",
|
||||
"SpatialUpsample2x",
|
||||
),
|
||||
decoder_temporal_upsample: Tuple[str] = ("", "", "TimeUpsample2x", "TimeUpsample2x"),
|
||||
decoder_mid_resnet: str = "ResnetBlock3D",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.tile_sample_min_size = 256
|
||||
self.tile_sample_min_size_t = 65
|
||||
self.tile_latent_min_size = int(self.tile_sample_min_size / (2 ** (len(hidden_size_mult) - 1)))
|
||||
t_down_ratio = [i for i in encoder_temporal_downsample if len(i) > 0]
|
||||
self.tile_latent_min_size_t = int((self.tile_sample_min_size_t - 1) / (2 ** len(t_down_ratio))) + 1
|
||||
self.tile_overlap_factor = 0.25
|
||||
self.use_tiling = False
|
||||
|
||||
self.learning_rate = lr
|
||||
self.lr_g_factor = 1.0
|
||||
|
||||
self.loss = resolve_str_to_obj(loss_type, append=False)(**loss_params)
|
||||
|
||||
self.encoder = Encoder(
|
||||
z_channels=z_channels,
|
||||
hidden_size=hidden_size,
|
||||
hidden_size_mult=hidden_size_mult,
|
||||
attn_resolutions=attn_resolutions,
|
||||
conv_in=encoder_conv_in,
|
||||
conv_out=encoder_conv_out,
|
||||
attention=encoder_attention,
|
||||
resnet_blocks=encoder_resnet_blocks,
|
||||
spatial_downsample=encoder_spatial_downsample,
|
||||
temporal_downsample=encoder_temporal_downsample,
|
||||
mid_resnet=encoder_mid_resnet,
|
||||
dropout=dropout,
|
||||
resolution=resolution,
|
||||
num_res_blocks=num_res_blocks,
|
||||
double_z=double_z,
|
||||
)
|
||||
|
||||
self.decoder = Decoder(
|
||||
z_channels=z_channels,
|
||||
hidden_size=hidden_size,
|
||||
hidden_size_mult=hidden_size_mult,
|
||||
attn_resolutions=attn_resolutions,
|
||||
conv_in=decoder_conv_in,
|
||||
conv_out=decoder_conv_out,
|
||||
attention=decoder_attention,
|
||||
resnet_blocks=decoder_resnet_blocks,
|
||||
spatial_upsample=decoder_spatial_upsample,
|
||||
temporal_upsample=decoder_temporal_upsample,
|
||||
mid_resnet=decoder_mid_resnet,
|
||||
dropout=dropout,
|
||||
resolution=resolution,
|
||||
num_res_blocks=num_res_blocks,
|
||||
)
|
||||
|
||||
quant_conv_cls = resolve_str_to_obj(q_conv)
|
||||
self.quant_conv = quant_conv_cls(2 * z_channels, 2 * embed_dim, 1)
|
||||
self.post_quant_conv = quant_conv_cls(embed_dim, z_channels, 1)
|
||||
if hasattr(self.loss, "discriminator"):
|
||||
self.automatic_optimization = False
|
||||
|
||||
def encode(self, x):
|
||||
if self.use_tiling and (
|
||||
x.shape[-1] > self.tile_sample_min_size
|
||||
or x.shape[-2] > self.tile_sample_min_size
|
||||
or x.shape[-3] > self.tile_sample_min_size_t
|
||||
):
|
||||
return self.tiled_encode(x)
|
||||
h = self.encoder(x)
|
||||
moments = self.quant_conv(h)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
return posterior
|
||||
|
||||
def decode(self, z):
|
||||
if self.use_tiling and (
|
||||
z.shape[-1] > self.tile_latent_min_size
|
||||
or z.shape[-2] > self.tile_latent_min_size
|
||||
or z.shape[-3] > self.tile_latent_min_size_t
|
||||
):
|
||||
return self.tiled_decode(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.to(memory_format=torch.contiguous_format).float()
|
||||
return x
|
||||
|
||||
def training_step(self, batch, batch_idx):
|
||||
if hasattr(self.loss, "discriminator"):
|
||||
return self._training_step_gan(batch, batch_idx=batch_idx)
|
||||
else:
|
||||
return self._training_step(batch, batch_idx=batch_idx)
|
||||
|
||||
def _training_step(self, batch, batch_idx):
|
||||
inputs = self.get_input(batch, "video")
|
||||
reconstructions, posterior = self(inputs)
|
||||
aeloss, log_dict_ae = self.loss(
|
||||
inputs,
|
||||
reconstructions,
|
||||
posterior,
|
||||
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
|
||||
|
||||
def _training_step_gan(self, batch, batch_idx):
|
||||
inputs = self.get_input(batch, "video")
|
||||
reconstructions, posterior = self(inputs)
|
||||
opt1, opt2 = self.optimizers()
|
||||
|
||||
# ---- AE Loss ----
|
||||
aeloss, log_dict_ae = self.loss(
|
||||
inputs,
|
||||
reconstructions,
|
||||
posterior,
|
||||
0,
|
||||
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,
|
||||
)
|
||||
opt1.zero_grad()
|
||||
self.manual_backward(aeloss)
|
||||
self.clip_gradients(opt1, gradient_clip_val=1, gradient_clip_algorithm="norm")
|
||||
opt1.step()
|
||||
# ---- GAN Loss ----
|
||||
discloss, log_dict_disc = self.loss(
|
||||
inputs,
|
||||
reconstructions,
|
||||
posterior,
|
||||
1,
|
||||
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,
|
||||
)
|
||||
opt2.zero_grad()
|
||||
self.manual_backward(discloss)
|
||||
self.clip_gradients(opt2, gradient_clip_val=1, gradient_clip_algorithm="norm")
|
||||
opt2.step()
|
||||
self.log_dict(
|
||||
{**log_dict_ae, **log_dict_disc},
|
||||
prog_bar=False,
|
||||
logger=True,
|
||||
on_step=True,
|
||||
on_epoch=False,
|
||||
)
|
||||
|
||||
def configure_optimizers(self):
|
||||
from itertools import chain
|
||||
|
||||
lr = self.learning_rate
|
||||
modules_to_train = [
|
||||
self.encoder.named_parameters(),
|
||||
self.decoder.named_parameters(),
|
||||
self.post_quant_conv.named_parameters(),
|
||||
self.quant_conv.named_parameters(),
|
||||
]
|
||||
params_with_time = []
|
||||
params_without_time = []
|
||||
for name, param in chain(*modules_to_train):
|
||||
if "time" in name:
|
||||
params_with_time.append(param)
|
||||
else:
|
||||
params_without_time.append(param)
|
||||
optimizers = []
|
||||
opt_ae = torch.optim.Adam(
|
||||
[
|
||||
{"params": params_with_time, "lr": lr},
|
||||
{"params": params_without_time, "lr": lr},
|
||||
],
|
||||
lr=lr,
|
||||
betas=(0.5, 0.9),
|
||||
)
|
||||
optimizers.append(opt_ae)
|
||||
|
||||
if hasattr(self.loss, "discriminator"):
|
||||
opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(), lr=lr, betas=(0.5, 0.9))
|
||||
optimizers.append(opt_disc)
|
||||
|
||||
return optimizers, []
|
||||
|
||||
def get_last_layer(self):
|
||||
if hasattr(self.decoder.conv_out, "conv"):
|
||||
return self.decoder.conv_out.conv.weight
|
||||
else:
|
||||
return self.decoder.conv_out.weight
|
||||
|
||||
def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[3], b.shape[3], blend_extent)
|
||||
for y in range(blend_extent):
|
||||
b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * (
|
||||
y / blend_extent
|
||||
)
|
||||
return b
|
||||
|
||||
def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[4], b.shape[4], blend_extent)
|
||||
for x in range(blend_extent):
|
||||
b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * (
|
||||
x / blend_extent
|
||||
)
|
||||
return b
|
||||
|
||||
def tiled_encode(self, x):
|
||||
t = x.shape[2]
|
||||
t_chunk_idx = [i for i in range(0, t, self.tile_sample_min_size_t - 1)]
|
||||
if len(t_chunk_idx) == 1 and t_chunk_idx[0] == 0:
|
||||
t_chunk_start_end = [[0, t]]
|
||||
else:
|
||||
t_chunk_start_end = [[t_chunk_idx[i], t_chunk_idx[i + 1] + 1] for i in range(len(t_chunk_idx) - 1)]
|
||||
if t_chunk_start_end[-1][-1] > t:
|
||||
t_chunk_start_end[-1][-1] = t
|
||||
elif t_chunk_start_end[-1][-1] < t:
|
||||
last_start_end = [t_chunk_idx[-1], t]
|
||||
t_chunk_start_end.append(last_start_end)
|
||||
moments = []
|
||||
for idx, (start, end) in enumerate(t_chunk_start_end):
|
||||
chunk_x = x[:, :, start:end]
|
||||
if idx != 0:
|
||||
moment = self.tiled_encode2d(chunk_x, return_moments=True)[:, :, 1:]
|
||||
else:
|
||||
moment = self.tiled_encode2d(chunk_x, return_moments=True)
|
||||
moments.append(moment)
|
||||
moments = torch.cat(moments, dim=2)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
return posterior
|
||||
|
||||
def tiled_decode(self, x):
|
||||
t = x.shape[2]
|
||||
t_chunk_idx = [i for i in range(0, t, self.tile_latent_min_size_t - 1)]
|
||||
if len(t_chunk_idx) == 1 and t_chunk_idx[0] == 0:
|
||||
t_chunk_start_end = [[0, t]]
|
||||
else:
|
||||
t_chunk_start_end = [[t_chunk_idx[i], t_chunk_idx[i + 1] + 1] for i in range(len(t_chunk_idx) - 1)]
|
||||
if t_chunk_start_end[-1][-1] > t:
|
||||
t_chunk_start_end[-1][-1] = t
|
||||
elif t_chunk_start_end[-1][-1] < t:
|
||||
last_start_end = [t_chunk_idx[-1], t]
|
||||
t_chunk_start_end.append(last_start_end)
|
||||
dec_ = []
|
||||
for idx, (start, end) in enumerate(t_chunk_start_end):
|
||||
chunk_x = x[:, :, start:end]
|
||||
if idx != 0:
|
||||
dec = self.tiled_decode2d(chunk_x)[:, :, 1:]
|
||||
else:
|
||||
dec = self.tiled_decode2d(chunk_x)
|
||||
dec_.append(dec)
|
||||
dec_ = torch.cat(dec_, dim=2)
|
||||
return dec_
|
||||
|
||||
def tiled_encode2d(self, x, return_moments=False):
|
||||
overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
|
||||
row_limit = self.tile_latent_min_size - blend_extent
|
||||
|
||||
# Split the image into 512x512 tiles and encode them separately.
|
||||
rows = []
|
||||
for i in range(0, x.shape[3], overlap_size):
|
||||
row = []
|
||||
for j in range(0, x.shape[4], overlap_size):
|
||||
tile = x[
|
||||
:,
|
||||
:,
|
||||
:,
|
||||
i : i + self.tile_sample_min_size,
|
||||
j : j + self.tile_sample_min_size,
|
||||
]
|
||||
tile = self.encoder(tile)
|
||||
tile = self.quant_conv(tile)
|
||||
row.append(tile)
|
||||
rows.append(row)
|
||||
result_rows = []
|
||||
for i, row in enumerate(rows):
|
||||
result_row = []
|
||||
for j, tile in enumerate(row):
|
||||
# blend the above tile and the left tile
|
||||
# to the current tile and add the current tile to the result row
|
||||
if i > 0:
|
||||
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
|
||||
if j > 0:
|
||||
tile = self.blend_h(row[j - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :, :row_limit, :row_limit])
|
||||
result_rows.append(torch.cat(result_row, dim=4))
|
||||
|
||||
moments = torch.cat(result_rows, dim=3)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
if return_moments:
|
||||
return moments
|
||||
return posterior
|
||||
|
||||
def tiled_decode2d(self, z):
|
||||
overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
|
||||
row_limit = self.tile_sample_min_size - blend_extent
|
||||
|
||||
# Split z into overlapping 64x64 tiles and decode them separately.
|
||||
# The tiles have an overlap to avoid seams between tiles.
|
||||
rows = []
|
||||
for i in range(0, z.shape[3], overlap_size):
|
||||
row = []
|
||||
for j in range(0, z.shape[4], overlap_size):
|
||||
tile = z[
|
||||
:,
|
||||
:,
|
||||
:,
|
||||
i : i + self.tile_latent_min_size,
|
||||
j : j + self.tile_latent_min_size,
|
||||
]
|
||||
tile = self.post_quant_conv(tile)
|
||||
decoded = self.decoder(tile)
|
||||
row.append(decoded)
|
||||
rows.append(row)
|
||||
result_rows = []
|
||||
for i, row in enumerate(rows):
|
||||
result_row = []
|
||||
for j, tile in enumerate(row):
|
||||
# blend the above tile and the left tile
|
||||
# to the current tile and add the current tile to the result row
|
||||
if i > 0:
|
||||
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
|
||||
if j > 0:
|
||||
tile = self.blend_h(row[j - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :, :row_limit, :row_limit])
|
||||
result_rows.append(torch.cat(result_row, dim=4))
|
||||
|
||||
dec = torch.cat(result_rows, dim=3)
|
||||
return dec
|
||||
|
||||
def enable_tiling(self, use_tiling: bool = True):
|
||||
self.use_tiling = use_tiling
|
||||
|
||||
def disable_tiling(self):
|
||||
self.enable_tiling(False)
|
||||
|
||||
def init_from_ckpt(self, path, ignore_keys=list(), remove_loss=False):
|
||||
sd = torch.load(path, map_location="cpu")
|
||||
print("init from " + path)
|
||||
if "state_dict" in sd:
|
||||
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]
|
||||
self.load_state_dict(sd, strict=False)
|
||||
|
||||
def validation_step(self, batch, batch_idx):
|
||||
inputs = self.get_input(batch, "video")
|
||||
latents = self.encode(inputs).sample()
|
||||
video_recon = self.decode(latents)
|
||||
for idx in range(len(video_recon)):
|
||||
self.logger.log_video(f"recon {batch_idx} {idx}", [tensor_to_video(video_recon[idx])], fps=[10])
|
||||
|
||||
|
||||
class CausalVAEModelWrapper(nn.Module):
|
||||
def __init__(self, model_path, subfolder=None, cache_dir=None, **kwargs):
|
||||
super(CausalVAEModelWrapper, self).__init__()
|
||||
# if os.path.exists(ckpt):
|
||||
# self.vae = CausalVAEModel.load_from_checkpoint(ckpt)
|
||||
self.vae = CausalVAEModel.from_pretrained(model_path, subfolder=subfolder, cache_dir=cache_dir, **kwargs)
|
||||
|
||||
def encode(self, x): # b c t h w
|
||||
# x = self.vae.encode(x).sample()
|
||||
x = self.vae.encode(x).sample().mul_(0.18215)
|
||||
return x
|
||||
|
||||
def decode(self, x):
|
||||
# x = self.vae.decode(x)
|
||||
x = self.vae.decode(x / 0.18215)
|
||||
x = rearrange(x, "b c t h w -> b t c h w").contiguous()
|
||||
return x
|
||||
|
||||
def dtype(self):
|
||||
return self.vae.dtype
|
||||
|
||||
#
|
||||
# def device(self):
|
||||
# return self.vae.device
|
||||
|
||||
|
||||
videobase_ae_stride = {
|
||||
"CausalVAEModel_4x8x8": [4, 8, 8],
|
||||
}
|
||||
|
||||
videobase_ae_channel = {
|
||||
"CausalVAEModel_4x8x8": 4,
|
||||
}
|
||||
|
||||
videobase_ae = {
|
||||
"CausalVAEModel_4x8x8": CausalVAEModelWrapper,
|
||||
}
|
||||
|
||||
|
||||
ae_stride_config = {}
|
||||
ae_stride_config.update(videobase_ae_stride)
|
||||
|
||||
ae_channel_config = {}
|
||||
ae_channel_config.update(videobase_ae_channel)
|
||||
|
||||
|
||||
def getae_wrapper(ae):
|
||||
"""deprecation"""
|
||||
ae = videobase_ae.get(ae, None)
|
||||
assert ae is not None
|
||||
return ae
|
||||
Executable
+2743
File diff suppressed because it is too large
Load Diff
Executable
+677
@@ -0,0 +1,677 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
import functools
|
||||
import hashlib
|
||||
import os
|
||||
from collections import namedtuple
|
||||
|
||||
import requests
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from torch import nn
|
||||
from torchvision import models
|
||||
from tqdm import tqdm
|
||||
|
||||
from opendit.models.opensora_plan.modules.normalize import ActNorm
|
||||
|
||||
URL_MAP = {"vgg_lpips": "https://heibox.uni-heidelberg.de/f/607503859c864bc1b30b/?dl=1"}
|
||||
|
||||
CKPT_MAP = {"vgg_lpips": "vgg.pth"}
|
||||
|
||||
MD5_MAP = {"vgg_lpips": "d507d7349b931f0638a25a48a722f98a"}
|
||||
|
||||
|
||||
def download(url, local_path, chunk_size=1024):
|
||||
os.makedirs(os.path.split(local_path)[0], exist_ok=True)
|
||||
with requests.get(url, stream=True) as r:
|
||||
total_size = int(r.headers.get("content-length", 0))
|
||||
with tqdm(total=total_size, unit="B", unit_scale=True) as pbar:
|
||||
with open(local_path, "wb") as f:
|
||||
for data in r.iter_content(chunk_size=chunk_size):
|
||||
if data:
|
||||
f.write(data)
|
||||
pbar.update(chunk_size)
|
||||
|
||||
|
||||
def md5_hash(path):
|
||||
with open(path, "rb") as f:
|
||||
content = f.read()
|
||||
return hashlib.md5(content).hexdigest()
|
||||
|
||||
|
||||
def get_ckpt_path(name, root, check=False):
|
||||
assert name in URL_MAP
|
||||
path = os.path.join(root, CKPT_MAP[name])
|
||||
if not os.path.exists(path) or (check and not md5_hash(path) == MD5_MAP[name]):
|
||||
print("Downloading {} model from {} to {}".format(name, URL_MAP[name], path))
|
||||
download(URL_MAP[name], path)
|
||||
md5 = md5_hash(path)
|
||||
assert md5 == MD5_MAP[name], md5
|
||||
return path
|
||||
|
||||
|
||||
class LPIPS(nn.Module):
|
||||
# Learned perceptual metric
|
||||
def __init__(self, use_dropout=True):
|
||||
super().__init__()
|
||||
self.scaling_layer = ScalingLayer()
|
||||
self.chns = [64, 128, 256, 512, 512] # vg16 features
|
||||
self.net = vgg16(pretrained=True, requires_grad=False)
|
||||
self.lin0 = NetLinLayer(self.chns[0], use_dropout=use_dropout)
|
||||
self.lin1 = NetLinLayer(self.chns[1], use_dropout=use_dropout)
|
||||
self.lin2 = NetLinLayer(self.chns[2], use_dropout=use_dropout)
|
||||
self.lin3 = NetLinLayer(self.chns[3], use_dropout=use_dropout)
|
||||
self.lin4 = NetLinLayer(self.chns[4], use_dropout=use_dropout)
|
||||
self.load_from_pretrained()
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def load_from_pretrained(self, name="vgg_lpips"):
|
||||
ckpt = get_ckpt_path(name, "taming/modules/autoencoder/lpips")
|
||||
self.load_state_dict(torch.load(ckpt, map_location=torch.device("cpu")), strict=False)
|
||||
print("loaded pretrained LPIPS loss from {}".format(ckpt))
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, name="vgg_lpips"):
|
||||
if name != "vgg_lpips":
|
||||
raise NotImplementedError
|
||||
model = cls()
|
||||
ckpt = get_ckpt_path(name)
|
||||
model.load_state_dict(torch.load(ckpt, map_location=torch.device("cpu")), strict=False)
|
||||
return model
|
||||
|
||||
def forward(self, input, target):
|
||||
in0_input, in1_input = (self.scaling_layer(input), self.scaling_layer(target))
|
||||
outs0, outs1 = self.net(in0_input), self.net(in1_input)
|
||||
feats0, feats1, diffs = {}, {}, {}
|
||||
lins = [self.lin0, self.lin1, self.lin2, self.lin3, self.lin4]
|
||||
for kk in range(len(self.chns)):
|
||||
feats0[kk], feats1[kk] = normalize_tensor(outs0[kk]), normalize_tensor(outs1[kk])
|
||||
diffs[kk] = (feats0[kk] - feats1[kk]) ** 2
|
||||
|
||||
res = [spatial_average(lins[kk].model(diffs[kk]), keepdim=True) for kk in range(len(self.chns))]
|
||||
val = res[0]
|
||||
for l in range(1, len(self.chns)):
|
||||
val += res[l]
|
||||
return val
|
||||
|
||||
|
||||
class ScalingLayer(nn.Module):
|
||||
def __init__(self):
|
||||
super(ScalingLayer, self).__init__()
|
||||
self.register_buffer("shift", torch.Tensor([-0.030, -0.088, -0.188])[None, :, None, None])
|
||||
self.register_buffer("scale", torch.Tensor([0.458, 0.448, 0.450])[None, :, None, None])
|
||||
|
||||
def forward(self, inp):
|
||||
return (inp - self.shift) / self.scale
|
||||
|
||||
|
||||
class NetLinLayer(nn.Module):
|
||||
"""A single linear layer which does a 1x1 conv"""
|
||||
|
||||
def __init__(self, chn_in, chn_out=1, use_dropout=False):
|
||||
super(NetLinLayer, self).__init__()
|
||||
layers = (
|
||||
[
|
||||
nn.Dropout(),
|
||||
]
|
||||
if (use_dropout)
|
||||
else []
|
||||
)
|
||||
layers += [
|
||||
nn.Conv2d(chn_in, chn_out, 1, stride=1, padding=0, bias=False),
|
||||
]
|
||||
self.model = nn.Sequential(*layers)
|
||||
|
||||
|
||||
class vgg16(torch.nn.Module):
|
||||
def __init__(self, requires_grad=False, pretrained=True):
|
||||
super(vgg16, self).__init__()
|
||||
vgg_pretrained_features = models.vgg16(pretrained=pretrained).features
|
||||
self.slice1 = torch.nn.Sequential()
|
||||
self.slice2 = torch.nn.Sequential()
|
||||
self.slice3 = torch.nn.Sequential()
|
||||
self.slice4 = torch.nn.Sequential()
|
||||
self.slice5 = torch.nn.Sequential()
|
||||
self.N_slices = 5
|
||||
for x in range(4):
|
||||
self.slice1.add_module(str(x), vgg_pretrained_features[x])
|
||||
for x in range(4, 9):
|
||||
self.slice2.add_module(str(x), vgg_pretrained_features[x])
|
||||
for x in range(9, 16):
|
||||
self.slice3.add_module(str(x), vgg_pretrained_features[x])
|
||||
for x in range(16, 23):
|
||||
self.slice4.add_module(str(x), vgg_pretrained_features[x])
|
||||
for x in range(23, 30):
|
||||
self.slice5.add_module(str(x), vgg_pretrained_features[x])
|
||||
if not requires_grad:
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, X):
|
||||
h = self.slice1(X)
|
||||
h_relu1_2 = h
|
||||
h = self.slice2(h)
|
||||
h_relu2_2 = h
|
||||
h = self.slice3(h)
|
||||
h_relu3_3 = h
|
||||
h = self.slice4(h)
|
||||
h_relu4_3 = h
|
||||
h = self.slice5(h)
|
||||
h_relu5_3 = h
|
||||
vgg_outputs = namedtuple("VggOutputs", ["relu1_2", "relu2_2", "relu3_3", "relu4_3", "relu5_3"])
|
||||
out = vgg_outputs(h_relu1_2, h_relu2_2, h_relu3_3, h_relu4_3, h_relu5_3)
|
||||
return out
|
||||
|
||||
|
||||
def normalize_tensor(x, eps=1e-10):
|
||||
norm_factor = torch.sqrt(torch.sum(x**2, dim=1, keepdim=True))
|
||||
return x / (norm_factor + eps)
|
||||
|
||||
|
||||
def spatial_average(x, keepdim=True):
|
||||
return x.mean([2, 3], keepdim=keepdim)
|
||||
|
||||
|
||||
def weights_init(m):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv") != -1:
|
||||
nn.init.normal_(m.weight.data, 0.0, 0.02)
|
||||
elif classname.find("BatchNorm") != -1:
|
||||
nn.init.normal_(m.weight.data, 1.0, 0.02)
|
||||
nn.init.constant_(m.bias.data, 0)
|
||||
|
||||
|
||||
def weights_init_conv(m):
|
||||
if hasattr(m, "conv"):
|
||||
m = m.conv
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv") != -1:
|
||||
nn.init.normal_(m.weight.data, 0.0, 0.02)
|
||||
elif classname.find("BatchNorm") != -1:
|
||||
nn.init.normal_(m.weight.data, 1.0, 0.02)
|
||||
nn.init.constant_(m.bias.data, 0)
|
||||
|
||||
|
||||
class NLayerDiscriminator(nn.Module):
|
||||
"""Defines a PatchGAN discriminator as in Pix2Pix
|
||||
--> see https://github.com/junyanz/pytorch-CycleGAN-and-pix2pix/blob/master/models/networks.py
|
||||
"""
|
||||
|
||||
def __init__(self, input_nc=3, ndf=64, n_layers=3, use_actnorm=False):
|
||||
"""Construct a PatchGAN discriminator
|
||||
Parameters:
|
||||
input_nc (int) -- the number of channels in input images
|
||||
ndf (int) -- the number of filters in the last conv layer
|
||||
n_layers (int) -- the number of conv layers in the discriminator
|
||||
norm_layer -- normalization layer
|
||||
"""
|
||||
super(NLayerDiscriminator, self).__init__()
|
||||
if not use_actnorm:
|
||||
norm_layer = nn.BatchNorm2d
|
||||
else:
|
||||
norm_layer = ActNorm
|
||||
if type(norm_layer) == functools.partial: # no need to use bias as BatchNorm2d has affine parameters
|
||||
use_bias = norm_layer.func != nn.BatchNorm2d
|
||||
else:
|
||||
use_bias = norm_layer != nn.BatchNorm2d
|
||||
|
||||
kw = 4
|
||||
padw = 1
|
||||
sequence = [nn.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw), nn.LeakyReLU(0.2, True)]
|
||||
nf_mult = 1
|
||||
nf_mult_prev = 1
|
||||
for n in range(1, n_layers): # gradually increase the number of filters
|
||||
nf_mult_prev = nf_mult
|
||||
nf_mult = min(2**n, 8)
|
||||
sequence += [
|
||||
nn.Conv2d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=2, padding=padw, bias=use_bias),
|
||||
norm_layer(ndf * nf_mult),
|
||||
nn.LeakyReLU(0.2, True),
|
||||
]
|
||||
|
||||
nf_mult_prev = nf_mult
|
||||
nf_mult = min(2**n_layers, 8)
|
||||
sequence += [
|
||||
nn.Conv2d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=1, padding=padw, bias=use_bias),
|
||||
norm_layer(ndf * nf_mult),
|
||||
nn.LeakyReLU(0.2, True),
|
||||
]
|
||||
|
||||
sequence += [
|
||||
nn.Conv2d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)
|
||||
] # output 1 channel prediction map
|
||||
self.main = nn.Sequential(*sequence)
|
||||
|
||||
def forward(self, input):
|
||||
"""Standard forward."""
|
||||
return self.main(input)
|
||||
|
||||
|
||||
class NLayerDiscriminator3D(nn.Module):
|
||||
"""Defines a 3D PatchGAN discriminator as in Pix2Pix but for 3D inputs."""
|
||||
|
||||
def __init__(self, input_nc=1, ndf=64, n_layers=3, use_actnorm=False):
|
||||
"""
|
||||
Construct a 3D PatchGAN discriminator
|
||||
|
||||
Parameters:
|
||||
input_nc (int) -- the number of channels in input volumes
|
||||
ndf (int) -- the number of filters in the last conv layer
|
||||
n_layers (int) -- the number of conv layers in the discriminator
|
||||
use_actnorm (bool) -- flag to use actnorm instead of batchnorm
|
||||
"""
|
||||
super(NLayerDiscriminator3D, self).__init__()
|
||||
if not use_actnorm:
|
||||
norm_layer = nn.BatchNorm3d
|
||||
else:
|
||||
raise NotImplementedError("Not implemented.")
|
||||
if type(norm_layer) == functools.partial:
|
||||
use_bias = norm_layer.func != nn.BatchNorm3d
|
||||
else:
|
||||
use_bias = norm_layer != nn.BatchNorm3d
|
||||
|
||||
kw = 3
|
||||
padw = 1
|
||||
sequence = [nn.Conv3d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw), nn.LeakyReLU(0.2, True)]
|
||||
nf_mult = 1
|
||||
nf_mult_prev = 1
|
||||
for n in range(1, n_layers): # gradually increase the number of filters
|
||||
nf_mult_prev = nf_mult
|
||||
nf_mult = min(2**n, 8)
|
||||
sequence += [
|
||||
nn.Conv3d(
|
||||
ndf * nf_mult_prev,
|
||||
ndf * nf_mult,
|
||||
kernel_size=(kw, kw, kw),
|
||||
stride=(2 if n == 1 else 1, 2, 2),
|
||||
padding=padw,
|
||||
bias=use_bias,
|
||||
),
|
||||
norm_layer(ndf * nf_mult),
|
||||
nn.LeakyReLU(0.2, True),
|
||||
]
|
||||
|
||||
nf_mult_prev = nf_mult
|
||||
nf_mult = min(2**n_layers, 8)
|
||||
sequence += [
|
||||
nn.Conv3d(
|
||||
ndf * nf_mult_prev, ndf * nf_mult, kernel_size=(kw, kw, kw), stride=1, padding=padw, bias=use_bias
|
||||
),
|
||||
norm_layer(ndf * nf_mult),
|
||||
nn.LeakyReLU(0.2, True),
|
||||
]
|
||||
|
||||
sequence += [
|
||||
nn.Conv3d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)
|
||||
] # output 1 channel prediction map
|
||||
self.main = nn.Sequential(*sequence)
|
||||
|
||||
def forward(self, input):
|
||||
"""Standard forward."""
|
||||
return self.main(input)
|
||||
|
||||
|
||||
def hinge_d_loss(logits_real, logits_fake):
|
||||
loss_real = torch.mean(F.relu(1.0 - logits_real))
|
||||
loss_fake = torch.mean(F.relu(1.0 + logits_fake))
|
||||
d_loss = 0.5 * (loss_real + loss_fake)
|
||||
return d_loss
|
||||
|
||||
|
||||
def vanilla_d_loss(logits_real, logits_fake):
|
||||
d_loss = 0.5 * (
|
||||
torch.mean(torch.nn.functional.softplus(-logits_real)) + torch.mean(torch.nn.functional.softplus(logits_fake))
|
||||
)
|
||||
return d_loss
|
||||
|
||||
|
||||
def hinge_d_loss_with_exemplar_weights(logits_real, logits_fake, weights):
|
||||
assert weights.shape[0] == logits_real.shape[0] == logits_fake.shape[0]
|
||||
loss_real = torch.mean(F.relu(1.0 - logits_real), dim=[1, 2, 3])
|
||||
loss_fake = torch.mean(F.relu(1.0 + logits_fake), dim=[1, 2, 3])
|
||||
loss_real = (weights * loss_real).sum() / weights.sum()
|
||||
loss_fake = (weights * loss_fake).sum() / weights.sum()
|
||||
d_loss = 0.5 * (loss_real + loss_fake)
|
||||
return d_loss
|
||||
|
||||
|
||||
def adopt_weight(weight, global_step, threshold=0, value=0.0):
|
||||
if global_step < threshold:
|
||||
weight = value
|
||||
return weight
|
||||
|
||||
|
||||
def measure_perplexity(predicted_indices, n_embed):
|
||||
# src: https://github.com/karpathy/deep-vector-quantization/blob/main/model.py
|
||||
# eval cluster perplexity. when perplexity == num_embeddings then all clusters are used exactly equally
|
||||
encodings = F.one_hot(predicted_indices, n_embed).float().reshape(-1, n_embed)
|
||||
avg_probs = encodings.mean(0)
|
||||
perplexity = (-(avg_probs * torch.log(avg_probs + 1e-10)).sum()).exp()
|
||||
cluster_use = torch.sum(avg_probs > 0)
|
||||
return perplexity, cluster_use
|
||||
|
||||
|
||||
def l1(x, y):
|
||||
return torch.abs(x - y)
|
||||
|
||||
|
||||
def l2(x, y):
|
||||
return torch.pow((x - y), 2)
|
||||
|
||||
|
||||
class LPIPSWithDiscriminator(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
disc_start,
|
||||
logvar_init=0.0,
|
||||
kl_weight=1.0,
|
||||
pixelloss_weight=1.0,
|
||||
perceptual_weight=1.0,
|
||||
# --- Discriminator Loss ---
|
||||
disc_num_layers=3,
|
||||
disc_in_channels=3,
|
||||
disc_factor=1.0,
|
||||
disc_weight=1.0,
|
||||
use_actnorm=False,
|
||||
disc_conditional=False,
|
||||
disc_loss="hinge",
|
||||
):
|
||||
super().__init__()
|
||||
assert disc_loss in ["hinge", "vanilla"]
|
||||
self.kl_weight = kl_weight
|
||||
self.pixel_weight = pixelloss_weight
|
||||
self.perceptual_loss = LPIPS().eval()
|
||||
self.perceptual_weight = perceptual_weight
|
||||
self.logvar = nn.Parameter(torch.ones(size=()) * logvar_init)
|
||||
|
||||
self.discriminator = NLayerDiscriminator(
|
||||
input_nc=disc_in_channels, n_layers=disc_num_layers, use_actnorm=use_actnorm
|
||||
).apply(weights_init)
|
||||
self.discriminator_iter_start = disc_start
|
||||
self.disc_loss = hinge_d_loss if disc_loss == "hinge" else vanilla_d_loss
|
||||
self.disc_factor = disc_factor
|
||||
self.discriminator_weight = disc_weight
|
||||
self.disc_conditional = disc_conditional
|
||||
|
||||
def calculate_adaptive_weight(self, nll_loss, g_loss, last_layer=None):
|
||||
if last_layer is not None:
|
||||
nll_grads = torch.autograd.grad(nll_loss, last_layer, retain_graph=True)[0]
|
||||
g_grads = torch.autograd.grad(g_loss, last_layer, retain_graph=True)[0]
|
||||
else:
|
||||
nll_grads = torch.autograd.grad(nll_loss, self.last_layer[0], retain_graph=True)[0]
|
||||
g_grads = torch.autograd.grad(g_loss, self.last_layer[0], retain_graph=True)[0]
|
||||
|
||||
d_weight = torch.norm(nll_grads) / (torch.norm(g_grads) + 1e-4)
|
||||
d_weight = torch.clamp(d_weight, 0.0, 1e4).detach()
|
||||
d_weight = d_weight * self.discriminator_weight
|
||||
return d_weight
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inputs,
|
||||
reconstructions,
|
||||
posteriors,
|
||||
optimizer_idx,
|
||||
global_step,
|
||||
split="train",
|
||||
weights=None,
|
||||
last_layer=None,
|
||||
cond=None,
|
||||
):
|
||||
inputs = rearrange(inputs, "b c t h w -> (b t) c h w").contiguous()
|
||||
reconstructions = rearrange(reconstructions, "b c t h w -> (b t) c h w").contiguous()
|
||||
rec_loss = torch.abs(inputs - reconstructions)
|
||||
if self.perceptual_weight > 0:
|
||||
p_loss = self.perceptual_loss(inputs, reconstructions)
|
||||
rec_loss = rec_loss + self.perceptual_weight * p_loss
|
||||
nll_loss = rec_loss / torch.exp(self.logvar) + self.logvar
|
||||
weighted_nll_loss = nll_loss
|
||||
if weights is not None:
|
||||
weighted_nll_loss = weights * nll_loss
|
||||
weighted_nll_loss = torch.sum(weighted_nll_loss) / weighted_nll_loss.shape[0]
|
||||
nll_loss = torch.sum(nll_loss) / nll_loss.shape[0]
|
||||
kl_loss = posteriors.kl()
|
||||
kl_loss = torch.sum(kl_loss) / kl_loss.shape[0]
|
||||
|
||||
# GAN Part
|
||||
if optimizer_idx == 0:
|
||||
# generator update
|
||||
if cond is None:
|
||||
assert not self.disc_conditional
|
||||
logits_fake = self.discriminator(reconstructions.contiguous())
|
||||
else:
|
||||
assert self.disc_conditional
|
||||
logits_fake = self.discriminator(torch.cat((reconstructions.contiguous(), cond), dim=1))
|
||||
g_loss = -torch.mean(logits_fake)
|
||||
|
||||
if self.disc_factor > 0.0:
|
||||
try:
|
||||
d_weight = self.calculate_adaptive_weight(nll_loss, g_loss, last_layer=last_layer)
|
||||
except RuntimeError:
|
||||
assert not self.training
|
||||
d_weight = torch.tensor(0.0)
|
||||
else:
|
||||
d_weight = torch.tensor(0.0)
|
||||
|
||||
disc_factor = adopt_weight(self.disc_factor, global_step, threshold=self.discriminator_iter_start)
|
||||
loss = weighted_nll_loss + self.kl_weight * kl_loss + d_weight * disc_factor * g_loss
|
||||
log = {
|
||||
"{}/total_loss".format(split): loss.clone().detach().mean(),
|
||||
"{}/logvar".format(split): self.logvar.detach(),
|
||||
"{}/kl_loss".format(split): kl_loss.detach().mean(),
|
||||
"{}/nll_loss".format(split): nll_loss.detach().mean(),
|
||||
"{}/rec_loss".format(split): rec_loss.detach().mean(),
|
||||
"{}/d_weight".format(split): d_weight.detach(),
|
||||
"{}/disc_factor".format(split): torch.tensor(disc_factor),
|
||||
"{}/g_loss".format(split): g_loss.detach().mean(),
|
||||
}
|
||||
return loss, log
|
||||
|
||||
if optimizer_idx == 1:
|
||||
if cond is None:
|
||||
logits_real = self.discriminator(inputs.contiguous().detach())
|
||||
logits_fake = self.discriminator(reconstructions.contiguous().detach())
|
||||
else:
|
||||
logits_real = self.discriminator(torch.cat((inputs.contiguous().detach(), cond), dim=1))
|
||||
logits_fake = self.discriminator(torch.cat((reconstructions.contiguous().detach(), cond), dim=1))
|
||||
|
||||
disc_factor = adopt_weight(self.disc_factor, global_step, threshold=self.discriminator_iter_start)
|
||||
d_loss = disc_factor * self.disc_loss(logits_real, logits_fake)
|
||||
|
||||
log = {
|
||||
"{}/disc_loss".format(split): d_loss.clone().detach().mean(),
|
||||
"{}/logits_real".format(split): logits_real.detach().mean(),
|
||||
"{}/logits_fake".format(split): logits_fake.detach().mean(),
|
||||
}
|
||||
return d_loss, log
|
||||
|
||||
|
||||
class LPIPSWithDiscriminator3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
disc_start,
|
||||
logvar_init=0.0,
|
||||
kl_weight=1.0,
|
||||
pixelloss_weight=1.0,
|
||||
perceptual_weight=1.0,
|
||||
# --- Discriminator Loss ---
|
||||
disc_num_layers=3,
|
||||
disc_in_channels=3,
|
||||
disc_factor=1.0,
|
||||
disc_weight=1.0,
|
||||
use_actnorm=False,
|
||||
disc_conditional=False,
|
||||
disc_loss="hinge",
|
||||
):
|
||||
super().__init__()
|
||||
assert disc_loss in ["hinge", "vanilla"]
|
||||
self.kl_weight = kl_weight
|
||||
self.pixel_weight = pixelloss_weight
|
||||
self.perceptual_loss = LPIPS().eval()
|
||||
self.perceptual_weight = perceptual_weight
|
||||
self.logvar = nn.Parameter(torch.ones(size=()) * logvar_init)
|
||||
|
||||
self.discriminator = NLayerDiscriminator3D(
|
||||
input_nc=disc_in_channels, n_layers=disc_num_layers, use_actnorm=use_actnorm
|
||||
).apply(weights_init)
|
||||
self.discriminator_iter_start = disc_start
|
||||
self.disc_loss = hinge_d_loss if disc_loss == "hinge" else vanilla_d_loss
|
||||
self.disc_factor = disc_factor
|
||||
self.discriminator_weight = disc_weight
|
||||
self.disc_conditional = disc_conditional
|
||||
|
||||
def calculate_adaptive_weight(self, nll_loss, g_loss, last_layer=None):
|
||||
if last_layer is not None:
|
||||
nll_grads = torch.autograd.grad(nll_loss, last_layer, retain_graph=True)[0]
|
||||
g_grads = torch.autograd.grad(g_loss, last_layer, retain_graph=True)[0]
|
||||
else:
|
||||
nll_grads = torch.autograd.grad(nll_loss, self.last_layer[0], retain_graph=True)[0]
|
||||
g_grads = torch.autograd.grad(g_loss, self.last_layer[0], retain_graph=True)[0]
|
||||
|
||||
d_weight = torch.norm(nll_grads) / (torch.norm(g_grads) + 1e-4)
|
||||
d_weight = torch.clamp(d_weight, 0.0, 1e4).detach()
|
||||
d_weight = d_weight * self.discriminator_weight
|
||||
return d_weight
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inputs,
|
||||
reconstructions,
|
||||
posteriors,
|
||||
optimizer_idx,
|
||||
global_step,
|
||||
split="train",
|
||||
weights=None,
|
||||
last_layer=None,
|
||||
cond=None,
|
||||
):
|
||||
t = inputs.shape[2]
|
||||
inputs = rearrange(inputs, "b c t h w -> (b t) c h w").contiguous()
|
||||
reconstructions = rearrange(reconstructions, "b c t h w -> (b t) c h w").contiguous()
|
||||
rec_loss = torch.abs(inputs - reconstructions)
|
||||
if self.perceptual_weight > 0:
|
||||
p_loss = self.perceptual_loss(inputs, reconstructions)
|
||||
rec_loss = rec_loss + self.perceptual_weight * p_loss
|
||||
nll_loss = rec_loss / torch.exp(self.logvar) + self.logvar
|
||||
weighted_nll_loss = nll_loss
|
||||
if weights is not None:
|
||||
weighted_nll_loss = weights * nll_loss
|
||||
weighted_nll_loss = torch.sum(weighted_nll_loss) / weighted_nll_loss.shape[0]
|
||||
nll_loss = torch.sum(nll_loss) / nll_loss.shape[0]
|
||||
kl_loss = posteriors.kl()
|
||||
kl_loss = torch.sum(kl_loss) / kl_loss.shape[0]
|
||||
inputs = rearrange(inputs, "(b t) c h w -> b c t h w", t=t).contiguous()
|
||||
reconstructions = rearrange(reconstructions, "(b t) c h w -> b c t h w", t=t).contiguous()
|
||||
# GAN Part
|
||||
if optimizer_idx == 0:
|
||||
# generator update
|
||||
if cond is None:
|
||||
assert not self.disc_conditional
|
||||
logits_fake = self.discriminator(reconstructions)
|
||||
else:
|
||||
assert self.disc_conditional
|
||||
logits_fake = self.discriminator(torch.cat((reconstructions, cond), dim=1))
|
||||
g_loss = -torch.mean(logits_fake)
|
||||
|
||||
if self.disc_factor > 0.0:
|
||||
try:
|
||||
d_weight = self.calculate_adaptive_weight(nll_loss, g_loss, last_layer=last_layer)
|
||||
except RuntimeError as e:
|
||||
assert not self.training, print(e)
|
||||
d_weight = torch.tensor(0.0)
|
||||
else:
|
||||
d_weight = torch.tensor(0.0)
|
||||
|
||||
disc_factor = adopt_weight(self.disc_factor, global_step, threshold=self.discriminator_iter_start)
|
||||
loss = weighted_nll_loss + self.kl_weight * kl_loss + d_weight * disc_factor * g_loss
|
||||
log = {
|
||||
"{}/total_loss".format(split): loss.clone().detach().mean(),
|
||||
"{}/logvar".format(split): self.logvar.detach(),
|
||||
"{}/kl_loss".format(split): kl_loss.detach().mean(),
|
||||
"{}/nll_loss".format(split): nll_loss.detach().mean(),
|
||||
"{}/rec_loss".format(split): rec_loss.detach().mean(),
|
||||
"{}/d_weight".format(split): d_weight.detach(),
|
||||
"{}/disc_factor".format(split): torch.tensor(disc_factor),
|
||||
"{}/g_loss".format(split): g_loss.detach().mean(),
|
||||
}
|
||||
return loss, log
|
||||
|
||||
if optimizer_idx == 1:
|
||||
if cond is None:
|
||||
logits_real = self.discriminator(inputs.contiguous().detach())
|
||||
logits_fake = self.discriminator(reconstructions.contiguous().detach())
|
||||
else:
|
||||
logits_real = self.discriminator(torch.cat((inputs.contiguous().detach(), cond), dim=1))
|
||||
logits_fake = self.discriminator(torch.cat((reconstructions.contiguous().detach(), cond), dim=1))
|
||||
|
||||
disc_factor = adopt_weight(self.disc_factor, global_step, threshold=self.discriminator_iter_start)
|
||||
d_loss = disc_factor * self.disc_loss(logits_real, logits_fake)
|
||||
|
||||
log = {
|
||||
"{}/disc_loss".format(split): d_loss.clone().detach().mean(),
|
||||
"{}/logits_real".format(split): logits_real.detach().mean(),
|
||||
"{}/logits_fake".format(split): logits_fake.detach().mean(),
|
||||
}
|
||||
return d_loss, log
|
||||
|
||||
|
||||
class SimpleLPIPS(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
logvar_init=0.0,
|
||||
kl_weight=1.0,
|
||||
pixelloss_weight=1.0,
|
||||
perceptual_weight=1.0,
|
||||
disc_loss="hinge",
|
||||
):
|
||||
super().__init__()
|
||||
assert disc_loss in ["hinge", "vanilla"]
|
||||
self.kl_weight = kl_weight
|
||||
self.pixel_weight = pixelloss_weight
|
||||
self.perceptual_loss = LPIPS().eval()
|
||||
self.perceptual_weight = perceptual_weight
|
||||
self.logvar = nn.Parameter(torch.ones(size=()) * logvar_init)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inputs,
|
||||
reconstructions,
|
||||
posteriors,
|
||||
split="train",
|
||||
weights=None,
|
||||
):
|
||||
inputs = rearrange(inputs, "b c t h w -> (b t) c h w").contiguous()
|
||||
reconstructions = rearrange(reconstructions, "b c t h w -> (b t) c h w").contiguous()
|
||||
rec_loss = torch.abs(inputs - reconstructions)
|
||||
if self.perceptual_weight > 0:
|
||||
p_loss = self.perceptual_loss(inputs, reconstructions)
|
||||
rec_loss = rec_loss + self.perceptual_weight * p_loss
|
||||
nll_loss = rec_loss / torch.exp(self.logvar) + self.logvar
|
||||
weighted_nll_loss = nll_loss
|
||||
if weights is not None:
|
||||
weighted_nll_loss = weights * nll_loss
|
||||
weighted_nll_loss = torch.sum(weighted_nll_loss) / weighted_nll_loss.shape[0]
|
||||
nll_loss = torch.sum(nll_loss) / nll_loss.shape[0]
|
||||
kl_loss = posteriors.kl()
|
||||
kl_loss = torch.sum(kl_loss) / kl_loss.shape[0]
|
||||
loss = weighted_nll_loss + self.kl_weight * kl_loss
|
||||
log = {
|
||||
"{}/total_loss".format(split): loss.clone().detach().mean(),
|
||||
"{}/logvar".format(split): self.logvar.detach(),
|
||||
"{}/kl_loss".format(split): kl_loss.detach().mean(),
|
||||
"{}/nll_loss".format(split): nll_loss.detach().mean(),
|
||||
"{}/rec_loss".format(split): rec_loss.detach().mean(),
|
||||
}
|
||||
if self.perceptual_weight > 0:
|
||||
log.update({"{}/p_loss".format(split): p_loss.detach().mean()})
|
||||
return loss, log
|
||||
+17
@@ -0,0 +1,17 @@
|
||||
from .attention import AttnBlock, AttnBlock3D, AttnBlock3DFix, LinAttnBlock, LinearAttention, TemporalAttnBlock
|
||||
from .block import Block
|
||||
from .conv import CausalConv3d, Conv2d
|
||||
from .normalize import GroupNorm, Normalize
|
||||
from .resnet_block import ResnetBlock2D, ResnetBlock3D
|
||||
from .updownsample import (
|
||||
Downsample,
|
||||
SpatialDownsample2x,
|
||||
SpatialUpsample2x,
|
||||
TimeDownsample2x,
|
||||
TimeDownsampleRes2x,
|
||||
TimeDownsampleResAdv2x,
|
||||
TimeUpsample2x,
|
||||
TimeUpsampleRes2x,
|
||||
TimeUpsampleResAdv2x,
|
||||
Upsample,
|
||||
)
|
||||
+227
@@ -0,0 +1,227 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from .block import Block
|
||||
from .conv import CausalConv3d
|
||||
from .normalize import Normalize
|
||||
from .ops import video_to_image
|
||||
|
||||
|
||||
class LinearAttention(Block):
|
||||
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)
|
||||
|
||||
|
||||
class AttnBlock3D(Block):
|
||||
"""Compatible with old versions, there are issues, use with caution."""
|
||||
|
||||
def __init__(self, in_channels):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = CausalConv3d(in_channels, in_channels, kernel_size=1, stride=1)
|
||||
self.k = CausalConv3d(in_channels, in_channels, kernel_size=1, stride=1)
|
||||
self.v = CausalConv3d(in_channels, in_channels, kernel_size=1, stride=1)
|
||||
self.proj_out = CausalConv3d(in_channels, in_channels, kernel_size=1, stride=1)
|
||||
|
||||
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, t, h, w = q.shape
|
||||
q = q.reshape(b * t, c, h * w)
|
||||
q = q.permute(0, 2, 1) # b,hw,c
|
||||
k = k.reshape(b * t, 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 * t, 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, t, h, w)
|
||||
|
||||
h_ = self.proj_out(h_)
|
||||
|
||||
return x + h_
|
||||
|
||||
|
||||
class AttnBlock3DFix(nn.Module):
|
||||
"""
|
||||
Thanks to https://github.com/PKU-YuanGroup/Open-Sora-Plan/pull/172.
|
||||
"""
|
||||
|
||||
def __init__(self, in_channels):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = CausalConv3d(in_channels, in_channels, kernel_size=1, stride=1)
|
||||
self.k = CausalConv3d(in_channels, in_channels, kernel_size=1, stride=1)
|
||||
self.v = CausalConv3d(in_channels, in_channels, kernel_size=1, stride=1)
|
||||
self.proj_out = CausalConv3d(in_channels, in_channels, kernel_size=1, stride=1)
|
||||
|
||||
def forward(self, x):
|
||||
h_ = x
|
||||
h_ = self.norm(h_)
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
# compute attention
|
||||
# q: (b c t h w) -> (b t c h w) -> (b*t c h*w) -> (b*t h*w c)
|
||||
b, c, t, h, w = q.shape
|
||||
q = q.permute(0, 2, 1, 3, 4)
|
||||
q = q.reshape(b * t, c, h * w)
|
||||
q = q.permute(0, 2, 1)
|
||||
|
||||
# k: (b c t h w) -> (b t c h w) -> (b*t c h*w)
|
||||
k = k.permute(0, 2, 1, 3, 4)
|
||||
k = k.reshape(b * t, c, h * w)
|
||||
|
||||
# w: (b*t hw hw)
|
||||
w_ = torch.bmm(q, k)
|
||||
w_ = w_ * (int(c) ** (-0.5))
|
||||
w_ = torch.nn.functional.softmax(w_, dim=2)
|
||||
|
||||
# attend to values
|
||||
# v: (b c t h w) -> (b t c h w) -> (bt c hw)
|
||||
# w_: (bt hw hw) -> (bt hw hw)
|
||||
v = v.permute(0, 2, 1, 3, 4)
|
||||
v = v.reshape(b * t, 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_: (b*t c hw) -> (b t c h w) -> (b c t h w)
|
||||
h_ = h_.reshape(b, t, c, h, w)
|
||||
h_ = h_.permute(0, 2, 1, 3, 4)
|
||||
|
||||
h_ = self.proj_out(h_)
|
||||
|
||||
return x + h_
|
||||
|
||||
|
||||
class AttnBlock(Block):
|
||||
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)
|
||||
|
||||
@video_to_image
|
||||
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 TemporalAttnBlock(Block):
|
||||
def __init__(self, in_channels):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = torch.nn.Conv3d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
self.k = torch.nn.Conv3d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
self.v = torch.nn.Conv3d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
self.proj_out = torch.nn.Conv3d(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, t, h, w = q.shape
|
||||
q = rearrange(q, "b c t h w -> (b h w) t c")
|
||||
k = rearrange(k, "b c t h w -> (b h w) c t")
|
||||
v = rearrange(v, "b c t h w -> (b h w) c t")
|
||||
w_ = torch.bmm(q, k)
|
||||
w_ = w_ * (int(c) ** (-0.5))
|
||||
w_ = torch.nn.functional.softmax(w_, dim=2)
|
||||
|
||||
# attend to values
|
||||
w_ = w_.permute(0, 2, 1)
|
||||
h_ = torch.bmm(v, w_)
|
||||
h_ = rearrange(h_, "(b h w) c t -> b c t h w", h=h, w=w)
|
||||
h_ = self.proj_out(h_)
|
||||
|
||||
return x + h_
|
||||
|
||||
|
||||
def make_attn(in_channels, attn_type="vanilla"):
|
||||
assert attn_type in ["vanilla", "linear", "none", "vanilla3D"], f"attn_type {attn_type} unknown"
|
||||
print(f"making attention of type '{attn_type}' with {in_channels} in_channels")
|
||||
print(attn_type)
|
||||
if attn_type == "vanilla":
|
||||
return AttnBlock(in_channels)
|
||||
elif attn_type == "vanilla3D":
|
||||
return AttnBlock3D(in_channels)
|
||||
elif attn_type == "none":
|
||||
return nn.Identity(in_channels)
|
||||
else:
|
||||
return LinAttnBlock(in_channels)
|
||||
+15
@@ -0,0 +1,15 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class Block(nn.Module):
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
Executable
+102
@@ -0,0 +1,102 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
from typing import Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from .ops import cast_tuple, video_to_image
|
||||
|
||||
|
||||
class Conv2d(nn.Conv2d):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
kernel_size: Union[int, Tuple[int]] = 3,
|
||||
stride: Union[int, Tuple[int]] = 1,
|
||||
padding: Union[str, int, Tuple[int]] = 0,
|
||||
dilation: Union[int, Tuple[int]] = 1,
|
||||
groups: int = 1,
|
||||
bias: bool = True,
|
||||
padding_mode: str = "zeros",
|
||||
device=None,
|
||||
dtype=None,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride,
|
||||
padding,
|
||||
dilation,
|
||||
groups,
|
||||
bias,
|
||||
padding_mode,
|
||||
device,
|
||||
dtype,
|
||||
)
|
||||
|
||||
@video_to_image
|
||||
def forward(self, x):
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
class CausalConv3d(nn.Module):
|
||||
def __init__(
|
||||
self, chan_in, chan_out, kernel_size: Union[int, Tuple[int, int, int]], init_method="random", **kwargs
|
||||
):
|
||||
super().__init__()
|
||||
self.kernel_size = cast_tuple(kernel_size, 3)
|
||||
self.time_kernel_size = self.kernel_size[0]
|
||||
self.chan_in = chan_in
|
||||
self.chan_out = chan_out
|
||||
stride = kwargs.pop("stride", 1)
|
||||
padding = kwargs.pop("padding", 0)
|
||||
padding = list(cast_tuple(padding, 3))
|
||||
padding[0] = 0
|
||||
stride = cast_tuple(stride, 3)
|
||||
self.conv = nn.Conv3d(chan_in, chan_out, self.kernel_size, stride=stride, padding=padding)
|
||||
self._init_weights(init_method)
|
||||
|
||||
def _init_weights(self, init_method):
|
||||
torch.tensor(self.kernel_size)
|
||||
if init_method == "avg":
|
||||
assert self.kernel_size[1] == 1 and self.kernel_size[2] == 1, "only support temporal up/down sample"
|
||||
assert self.chan_in == self.chan_out, "chan_in must be equal to chan_out"
|
||||
weight = torch.zeros((self.chan_out, self.chan_in, *self.kernel_size))
|
||||
|
||||
eyes = torch.concat(
|
||||
[
|
||||
torch.eye(self.chan_in).unsqueeze(-1) * 1 / 3,
|
||||
torch.eye(self.chan_in).unsqueeze(-1) * 1 / 3,
|
||||
torch.eye(self.chan_in).unsqueeze(-1) * 1 / 3,
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
weight[:, :, :, 0, 0] = eyes
|
||||
|
||||
self.conv.weight = nn.Parameter(
|
||||
weight,
|
||||
requires_grad=True,
|
||||
)
|
||||
elif init_method == "zero":
|
||||
self.conv.weight = nn.Parameter(
|
||||
torch.zeros((self.chan_out, self.chan_in, *self.kernel_size)),
|
||||
requires_grad=True,
|
||||
)
|
||||
if self.conv.bias is not None:
|
||||
nn.init.constant_(self.conv.bias, 0)
|
||||
|
||||
def forward(self, x):
|
||||
# 1 + 16 16 as video, 1 as image
|
||||
first_frame_pad = x[:, :, :1, :, :].repeat((1, 1, self.time_kernel_size - 1, 1, 1)) # b c t h w
|
||||
x = torch.concatenate((first_frame_pad, x), dim=2) # 3 + 16
|
||||
return self.conv(x)
|
||||
+98
@@ -0,0 +1,98 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from .block import Block
|
||||
|
||||
|
||||
class GroupNorm(Block):
|
||||
def __init__(self, num_channels, num_groups=32, eps=1e-6, *args, **kwargs) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
self.norm = torch.nn.GroupNorm(num_groups=num_groups, num_channels=num_channels, eps=1e-6, affine=True)
|
||||
|
||||
def forward(self, x):
|
||||
return self.norm(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 ActNorm(nn.Module):
|
||||
def __init__(self, num_features, logdet=False, affine=True, allow_reverse_init=False):
|
||||
assert affine
|
||||
super().__init__()
|
||||
self.logdet = logdet
|
||||
self.loc = nn.Parameter(torch.zeros(1, num_features, 1, 1))
|
||||
self.scale = nn.Parameter(torch.ones(1, num_features, 1, 1))
|
||||
self.allow_reverse_init = allow_reverse_init
|
||||
|
||||
self.register_buffer("initialized", torch.tensor(0, dtype=torch.uint8))
|
||||
|
||||
def initialize(self, input):
|
||||
with torch.no_grad():
|
||||
flatten = input.permute(1, 0, 2, 3).contiguous().view(input.shape[1], -1)
|
||||
mean = flatten.mean(1).unsqueeze(1).unsqueeze(2).unsqueeze(3).permute(1, 0, 2, 3)
|
||||
std = flatten.std(1).unsqueeze(1).unsqueeze(2).unsqueeze(3).permute(1, 0, 2, 3)
|
||||
|
||||
self.loc.data.copy_(-mean)
|
||||
self.scale.data.copy_(1 / (std + 1e-6))
|
||||
|
||||
def forward(self, input, reverse=False):
|
||||
if reverse:
|
||||
return self.reverse(input)
|
||||
if len(input.shape) == 2:
|
||||
input = input[:, :, None, None]
|
||||
squeeze = True
|
||||
else:
|
||||
squeeze = False
|
||||
|
||||
_, _, height, width = input.shape
|
||||
|
||||
if self.training and self.initialized.item() == 0:
|
||||
self.initialize(input)
|
||||
self.initialized.fill_(1)
|
||||
|
||||
h = self.scale * (input + self.loc)
|
||||
|
||||
if squeeze:
|
||||
h = h.squeeze(-1).squeeze(-1)
|
||||
|
||||
if self.logdet:
|
||||
log_abs = torch.log(torch.abs(self.scale))
|
||||
logdet = height * width * torch.sum(log_abs)
|
||||
logdet = logdet * torch.ones(input.shape[0]).to(input)
|
||||
return h, logdet
|
||||
|
||||
return h
|
||||
|
||||
def reverse(self, output):
|
||||
if self.training and self.initialized.item() == 0:
|
||||
if not self.allow_reverse_init:
|
||||
raise RuntimeError(
|
||||
"Initializing ActNorm in reverse direction is "
|
||||
"disabled by default. Use allow_reverse_init=True to enable."
|
||||
)
|
||||
else:
|
||||
self.initialize(output)
|
||||
self.initialized.fill_(1)
|
||||
|
||||
if len(output.shape) == 2:
|
||||
output = output[:, :, None, None]
|
||||
squeeze = True
|
||||
else:
|
||||
squeeze = False
|
||||
|
||||
h = output / self.scale - self.loc
|
||||
|
||||
if squeeze:
|
||||
h = h.squeeze(-1).squeeze(-1)
|
||||
return h
|
||||
Executable
+54
@@ -0,0 +1,54 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
def video_to_image(func):
|
||||
def wrapper(self, x, *args, **kwargs):
|
||||
if x.dim() == 5:
|
||||
t = x.shape[2]
|
||||
x = rearrange(x, "b c t h w -> (b t) c h w")
|
||||
x = func(self, x, *args, **kwargs)
|
||||
x = rearrange(x, "(b t) c h w -> b c t h w", t=t)
|
||||
return x
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
def nonlinearity(x):
|
||||
return x * torch.sigmoid(x)
|
||||
|
||||
|
||||
def cast_tuple(t, length=1):
|
||||
return t if isinstance(t, tuple) else ((t,) * length)
|
||||
|
||||
|
||||
def shift_dim(x, src_dim=-1, dest_dim=-1, make_contiguous=True):
|
||||
n_dims = len(x.shape)
|
||||
if src_dim < 0:
|
||||
src_dim = n_dims + src_dim
|
||||
if dest_dim < 0:
|
||||
dest_dim = n_dims + dest_dim
|
||||
assert 0 <= src_dim < n_dims and 0 <= dest_dim < n_dims
|
||||
dims = list(range(n_dims))
|
||||
del dims[src_dim]
|
||||
permutation = []
|
||||
ctr = 0
|
||||
for i in range(n_dims):
|
||||
if i == dest_dim:
|
||||
permutation.append(src_dim)
|
||||
else:
|
||||
permutation.append(dims[ctr])
|
||||
ctr += 1
|
||||
x = x.permute(permutation)
|
||||
if make_contiguous:
|
||||
x = x.contiguous()
|
||||
return x
|
||||
+111
@@ -0,0 +1,111 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .ops import shift_dim
|
||||
|
||||
|
||||
class Codebook(nn.Module):
|
||||
def __init__(self, n_codes, embedding_dim):
|
||||
super().__init__()
|
||||
self.register_buffer("embeddings", torch.randn(n_codes, embedding_dim))
|
||||
self.register_buffer("N", torch.zeros(n_codes))
|
||||
self.register_buffer("z_avg", self.embeddings.data.clone())
|
||||
|
||||
self.n_codes = n_codes
|
||||
self.embedding_dim = embedding_dim
|
||||
self._need_init = True
|
||||
|
||||
def _tile(self, x):
|
||||
d, ew = x.shape
|
||||
if d < self.n_codes:
|
||||
n_repeats = (self.n_codes + d - 1) // d
|
||||
std = 0.01 / np.sqrt(ew)
|
||||
x = x.repeat(n_repeats, 1)
|
||||
x = x + torch.randn_like(x) * std
|
||||
return x
|
||||
|
||||
def _init_embeddings(self, z):
|
||||
# z: [b, c, t, h, w]
|
||||
self._need_init = False
|
||||
flat_inputs = shift_dim(z, 1, -1).flatten(end_dim=-2)
|
||||
y = self._tile(flat_inputs)
|
||||
|
||||
y.shape[0]
|
||||
_k_rand = y[torch.randperm(y.shape[0])][: self.n_codes]
|
||||
if dist.is_initialized():
|
||||
dist.broadcast(_k_rand, 0)
|
||||
self.embeddings.data.copy_(_k_rand)
|
||||
self.z_avg.data.copy_(_k_rand)
|
||||
self.N.data.copy_(torch.ones(self.n_codes))
|
||||
|
||||
def forward(self, z):
|
||||
# z: [b, c, t, h, w]
|
||||
if self._need_init and self.training:
|
||||
self._init_embeddings(z)
|
||||
flat_inputs = shift_dim(z, 1, -1).flatten(end_dim=-2)
|
||||
distances = (
|
||||
(flat_inputs**2).sum(dim=1, keepdim=True)
|
||||
- 2 * flat_inputs @ self.embeddings.t()
|
||||
+ (self.embeddings.t() ** 2).sum(dim=0, keepdim=True)
|
||||
)
|
||||
|
||||
encoding_indices = torch.argmin(distances, dim=1)
|
||||
encode_onehot = F.one_hot(encoding_indices, self.n_codes).type_as(flat_inputs)
|
||||
encoding_indices = encoding_indices.view(z.shape[0], *z.shape[2:])
|
||||
|
||||
embeddings = F.embedding(encoding_indices, self.embeddings)
|
||||
embeddings = shift_dim(embeddings, -1, 1)
|
||||
|
||||
commitment_loss = 0.25 * F.mse_loss(z, embeddings.detach())
|
||||
|
||||
# EMA codebook update
|
||||
if self.training:
|
||||
n_total = encode_onehot.sum(dim=0)
|
||||
encode_sum = flat_inputs.t() @ encode_onehot
|
||||
if dist.is_initialized():
|
||||
dist.all_reduce(n_total)
|
||||
dist.all_reduce(encode_sum)
|
||||
|
||||
self.N.data.mul_(0.99).add_(n_total, alpha=0.01)
|
||||
self.z_avg.data.mul_(0.99).add_(encode_sum.t(), alpha=0.01)
|
||||
|
||||
n = self.N.sum()
|
||||
weights = (self.N + 1e-7) / (n + self.n_codes * 1e-7) * n
|
||||
encode_normalized = self.z_avg / weights.unsqueeze(1)
|
||||
self.embeddings.data.copy_(encode_normalized)
|
||||
|
||||
y = self._tile(flat_inputs)
|
||||
_k_rand = y[torch.randperm(y.shape[0])][: self.n_codes]
|
||||
if dist.is_initialized():
|
||||
dist.broadcast(_k_rand, 0)
|
||||
|
||||
usage = (self.N.view(self.n_codes, 1) >= 1).float()
|
||||
self.embeddings.data.mul_(usage).add_(_k_rand * (1 - usage))
|
||||
|
||||
embeddings_st = (embeddings - z).detach() + z
|
||||
|
||||
avg_probs = torch.mean(encode_onehot, dim=0)
|
||||
perplexity = torch.exp(-torch.sum(avg_probs * torch.log(avg_probs + 1e-10)))
|
||||
|
||||
return dict(
|
||||
embeddings=embeddings_st,
|
||||
encodings=encoding_indices,
|
||||
commitment_loss=commitment_loss,
|
||||
perplexity=perplexity,
|
||||
)
|
||||
|
||||
def dictionary_lookup(self, encodings):
|
||||
embeddings = F.embedding(encodings, self.embeddings)
|
||||
return embeddings
|
||||
+87
@@ -0,0 +1,87 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
import torch
|
||||
|
||||
from .block import Block
|
||||
from .conv import CausalConv3d
|
||||
from .normalize import Normalize
|
||||
from .ops import nonlinearity, video_to_image
|
||||
|
||||
|
||||
class ResnetBlock2D(Block):
|
||||
def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False, dropout):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = in_channels if out_channels is None else 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)
|
||||
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)
|
||||
|
||||
@video_to_image
|
||||
def forward(self, x):
|
||||
h = x
|
||||
h = self.norm1(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv1(h)
|
||||
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)
|
||||
x = x + h
|
||||
return x
|
||||
|
||||
|
||||
class ResnetBlock3D(Block):
|
||||
def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False, dropout):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = in_channels if out_channels is None else out_channels
|
||||
self.use_conv_shortcut = conv_shortcut
|
||||
|
||||
self.norm1 = Normalize(in_channels)
|
||||
self.conv1 = CausalConv3d(in_channels, out_channels, 3, padding=1)
|
||||
self.norm2 = Normalize(out_channels)
|
||||
self.dropout = torch.nn.Dropout(dropout)
|
||||
self.conv2 = CausalConv3d(out_channels, out_channels, 3, padding=1)
|
||||
if self.in_channels != self.out_channels:
|
||||
if self.use_conv_shortcut:
|
||||
self.conv_shortcut = CausalConv3d(in_channels, out_channels, 3, padding=1)
|
||||
else:
|
||||
self.nin_shortcut = CausalConv3d(in_channels, out_channels, 1, padding=0)
|
||||
|
||||
def forward(self, x):
|
||||
h = x
|
||||
h = self.norm1(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv1(h)
|
||||
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
|
||||
+215
@@ -0,0 +1,215 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
from typing import Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
from .attention import TemporalAttnBlock
|
||||
from .block import Block
|
||||
from .conv import CausalConv3d
|
||||
from .normalize import Normalize
|
||||
from .ops import cast_tuple, video_to_image
|
||||
from .resnet_block import ResnetBlock3D
|
||||
|
||||
|
||||
class Upsample(Block):
|
||||
def __init__(self, in_channels, out_channels):
|
||||
super().__init__()
|
||||
self.with_conv = True
|
||||
if self.with_conv:
|
||||
self.conv = torch.nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
@video_to_image
|
||||
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(Block):
|
||||
def __init__(self, in_channels, out_channels):
|
||||
super().__init__()
|
||||
self.with_conv = True
|
||||
if self.with_conv:
|
||||
# no asymmetric padding in torch conv, must do it ourselves
|
||||
self.conv = torch.nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=2, padding=0)
|
||||
|
||||
@video_to_image
|
||||
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 SpatialDownsample2x(Block):
|
||||
def __init__(
|
||||
self,
|
||||
chan_in,
|
||||
chan_out,
|
||||
kernel_size: Union[int, Tuple[int]] = (3, 3),
|
||||
stride: Union[int, Tuple[int]] = (2, 2),
|
||||
):
|
||||
super().__init__()
|
||||
kernel_size = cast_tuple(kernel_size, 2)
|
||||
stride = cast_tuple(stride, 2)
|
||||
self.chan_in = chan_in
|
||||
self.chan_out = chan_out
|
||||
self.kernel_size = kernel_size
|
||||
self.conv = CausalConv3d(self.chan_in, self.chan_out, (1,) + self.kernel_size, stride=(1,) + stride, padding=0)
|
||||
|
||||
def forward(self, x):
|
||||
pad = (0, 1, 0, 1, 0, 0)
|
||||
x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
|
||||
x = self.conv(x)
|
||||
return x
|
||||
|
||||
|
||||
class SpatialUpsample2x(Block):
|
||||
def __init__(
|
||||
self,
|
||||
chan_in,
|
||||
chan_out,
|
||||
kernel_size: Union[int, Tuple[int]] = (3, 3),
|
||||
stride: Union[int, Tuple[int]] = (1, 1),
|
||||
):
|
||||
super().__init__()
|
||||
self.chan_in = chan_in
|
||||
self.chan_out = chan_out
|
||||
self.kernel_size = kernel_size
|
||||
self.conv = CausalConv3d(self.chan_in, self.chan_out, (1,) + self.kernel_size, stride=(1,) + stride, padding=1)
|
||||
|
||||
def forward(self, x):
|
||||
t = x.shape[2]
|
||||
x = rearrange(x, "b c t h w -> b (c t) h w")
|
||||
x = F.interpolate(x, scale_factor=(2, 2), mode="nearest")
|
||||
x = rearrange(x, "b (c t) h w -> b c t h w", t=t)
|
||||
x = self.conv(x)
|
||||
return x
|
||||
|
||||
|
||||
class TimeDownsample2x(Block):
|
||||
def __init__(self, chan_in, chan_out, kernel_size: int = 3):
|
||||
super().__init__()
|
||||
self.kernel_size = kernel_size
|
||||
self.conv = nn.AvgPool3d((kernel_size, 1, 1), stride=(2, 1, 1))
|
||||
|
||||
def forward(self, x):
|
||||
first_frame_pad = x[:, :, :1, :, :].repeat((1, 1, self.kernel_size - 1, 1, 1))
|
||||
x = torch.concatenate((first_frame_pad, x), dim=2)
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class TimeUpsample2x(Block):
|
||||
def __init__(self, chan_in, chan_out):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, x):
|
||||
if x.size(2) > 1:
|
||||
x, x_ = x[:, :, :1], x[:, :, 1:]
|
||||
x_ = F.interpolate(x_, scale_factor=(2, 1, 1), mode="trilinear")
|
||||
x = torch.concat([x, x_], dim=2)
|
||||
return x
|
||||
|
||||
|
||||
class TimeDownsampleRes2x(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size: int = 3,
|
||||
mix_factor: float = 2.0,
|
||||
):
|
||||
super().__init__()
|
||||
self.kernel_size = cast_tuple(kernel_size, 3)
|
||||
self.avg_pool = nn.AvgPool3d((kernel_size, 1, 1), stride=(2, 1, 1))
|
||||
self.conv = nn.Conv3d(in_channels, out_channels, self.kernel_size, stride=(2, 1, 1), padding=(0, 1, 1))
|
||||
self.mix_factor = torch.nn.Parameter(torch.Tensor([mix_factor]))
|
||||
|
||||
def forward(self, x):
|
||||
alpha = torch.sigmoid(self.mix_factor)
|
||||
first_frame_pad = x[:, :, :1, :, :].repeat((1, 1, self.kernel_size[0] - 1, 1, 1))
|
||||
x = torch.concatenate((first_frame_pad, x), dim=2)
|
||||
return alpha * self.avg_pool(x) + (1 - alpha) * self.conv(x)
|
||||
|
||||
|
||||
class TimeUpsampleRes2x(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size: int = 3,
|
||||
mix_factor: float = 2.0,
|
||||
):
|
||||
super().__init__()
|
||||
self.conv = CausalConv3d(in_channels, out_channels, kernel_size, padding=1)
|
||||
self.mix_factor = torch.nn.Parameter(torch.Tensor([mix_factor]))
|
||||
|
||||
def forward(self, x):
|
||||
alpha = torch.sigmoid(self.mix_factor)
|
||||
if x.size(2) > 1:
|
||||
x, x_ = x[:, :, :1], x[:, :, 1:]
|
||||
x_ = F.interpolate(x_, scale_factor=(2, 1, 1), mode="trilinear")
|
||||
x = torch.concat([x, x_], dim=2)
|
||||
return alpha * x + (1 - alpha) * self.conv(x)
|
||||
|
||||
|
||||
class TimeDownsampleResAdv2x(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size: int = 3,
|
||||
mix_factor: float = 1.5,
|
||||
):
|
||||
super().__init__()
|
||||
self.kernel_size = cast_tuple(kernel_size, 3)
|
||||
self.avg_pool = nn.AvgPool3d((kernel_size, 1, 1), stride=(2, 1, 1))
|
||||
self.attn = TemporalAttnBlock(in_channels)
|
||||
self.res = ResnetBlock3D(in_channels=in_channels, out_channels=in_channels, dropout=0.0)
|
||||
self.conv = nn.Conv3d(in_channels, out_channels, self.kernel_size, stride=(2, 1, 1), padding=(0, 1, 1))
|
||||
self.mix_factor = torch.nn.Parameter(torch.Tensor([mix_factor]))
|
||||
|
||||
def forward(self, x):
|
||||
first_frame_pad = x[:, :, :1, :, :].repeat((1, 1, self.kernel_size[0] - 1, 1, 1))
|
||||
x = torch.concatenate((first_frame_pad, x), dim=2)
|
||||
alpha = torch.sigmoid(self.mix_factor)
|
||||
return alpha * self.avg_pool(x) + (1 - alpha) * self.conv(self.attn((self.res(x))))
|
||||
|
||||
|
||||
class TimeUpsampleResAdv2x(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size: int = 3,
|
||||
mix_factor: float = 1.5,
|
||||
):
|
||||
super().__init__()
|
||||
self.res = ResnetBlock3D(in_channels=in_channels, out_channels=in_channels, dropout=0.0)
|
||||
self.attn = TemporalAttnBlock(in_channels)
|
||||
self.norm = Normalize(in_channels=in_channels)
|
||||
self.conv = CausalConv3d(in_channels, out_channels, kernel_size, padding=1)
|
||||
self.mix_factor = torch.nn.Parameter(torch.Tensor([mix_factor]))
|
||||
|
||||
def forward(self, x):
|
||||
if x.size(2) > 1:
|
||||
x, x_ = x[:, :, :1], x[:, :, 1:]
|
||||
x_ = F.interpolate(x_, scale_factor=(2, 1, 1), mode="trilinear")
|
||||
x = torch.concat([x, x_], dim=2)
|
||||
alpha = torch.sigmoid(self.mix_factor)
|
||||
return alpha * x + (1 - alpha) * self.conv(self.attn(self.res(x)))
|
||||
Executable
+781
@@ -0,0 +1,781 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
import html
|
||||
import inspect
|
||||
import math
|
||||
import re
|
||||
import urllib.parse as ul
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from diffusers.models import AutoencoderKL, Transformer2DModel
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.schedulers import DPMSolverMultistepScheduler
|
||||
from diffusers.utils import (
|
||||
BACKENDS_MAPPING,
|
||||
BaseOutput,
|
||||
is_bs4_available,
|
||||
is_ftfy_available,
|
||||
logging,
|
||||
replace_example_docstring,
|
||||
)
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from transformers import T5EncoderModel, T5Tokenizer
|
||||
|
||||
from opendit.core.pab_mgr import get_diffusion_skip, get_diffusion_skip_timestep, skip_diffusion_timestep
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
if is_bs4_available():
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
if is_ftfy_available():
|
||||
import ftfy
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
Examples:
|
||||
```py
|
||||
>>> import torch
|
||||
>>> from diffusers import PixArtAlphaPipeline
|
||||
|
||||
>>> # You can replace the checkpoint id with "PixArt-alpha/PixArt-XL-2-512x512" too.
|
||||
>>> pipe = PixArtAlphaPipeline.from_pretrained("PixArt-alpha/PixArt-XL-2-1024-MS", torch_dtype=torch.float16)
|
||||
>>> # Enable memory optimizations.
|
||||
>>> pipe.enable_model_cpu_offload()
|
||||
|
||||
>>> prompt = "A small cactus with a happy face in the Sahara desert."
|
||||
>>> image = pipe(prompt).images[0]
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoPipelineOutput(BaseOutput):
|
||||
video: torch.Tensor
|
||||
|
||||
|
||||
class VideoGenPipeline(DiffusionPipeline):
|
||||
r"""
|
||||
Pipeline for text-to-image generation using PixArt-Alpha.
|
||||
|
||||
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.)
|
||||
|
||||
Args:
|
||||
vae ([`AutoencoderKL`]):
|
||||
Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.
|
||||
text_encoder ([`T5EncoderModel`]):
|
||||
Frozen text-encoder. PixArt-Alpha uses
|
||||
[T5](https://huggingface.co/docs/transformers/model_doc/t5#transformers.T5EncoderModel), specifically the
|
||||
[t5-v1_1-xxl](https://huggingface.co/PixArt-alpha/PixArt-alpha/tree/main/t5-v1_1-xxl) variant.
|
||||
tokenizer (`T5Tokenizer`):
|
||||
Tokenizer of class
|
||||
[T5Tokenizer](https://huggingface.co/docs/transformers/model_doc/t5#transformers.T5Tokenizer).
|
||||
transformer ([`Transformer2DModel`]):
|
||||
A text conditioned `Transformer2DModel` to denoise the encoded image latents.
|
||||
scheduler ([`SchedulerMixin`]):
|
||||
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
|
||||
"""
|
||||
bad_punct_regex = re.compile(
|
||||
r"[" + "#®•©™&@·º½¾¿¡§~" + "\)" + "\(" + "\]" + "\[" + "\}" + "\{" + "\|" + "\\" + "\/" + "\*" + r"]{1,}"
|
||||
) # noqa
|
||||
|
||||
_optional_components = ["tokenizer", "text_encoder"]
|
||||
model_cpu_offload_seq = "text_encoder->transformer->vae"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: T5Tokenizer,
|
||||
text_encoder: T5EncoderModel,
|
||||
vae: AutoencoderKL,
|
||||
transformer: Transformer2DModel,
|
||||
scheduler: DPMSolverMultistepScheduler,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
tokenizer=tokenizer, text_encoder=text_encoder, vae=vae, transformer=transformer, scheduler=scheduler
|
||||
)
|
||||
|
||||
# self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
|
||||
# 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:
|
||||
keep_index = mask.sum().item()
|
||||
return emb[:, :, :keep_index, :], keep_index # 1, 120, 4096 -> 1 7 4096
|
||||
else:
|
||||
masked_feature = emb * mask[:, None, :, None] # 1 120 4096
|
||||
return masked_feature, emb.shape[2]
|
||||
|
||||
# Adapted from diffusers.pipelines.deepfloyd_if.pipeline_if.encode_prompt
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
do_classifier_free_guidance: bool = True,
|
||||
negative_prompt: str = "",
|
||||
num_images_per_prompt: int = 1,
|
||||
device: Optional[torch.device] = None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
clean_caption: bool = False,
|
||||
mask_feature: bool = True,
|
||||
):
|
||||
r"""
|
||||
Encodes the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
prompt to be encoded
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt 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`). For
|
||||
PixArt-Alpha, this should be "".
|
||||
do_classifier_free_guidance (`bool`, *optional*, defaults to `True`):
|
||||
whether to use classifier free guidance or not
|
||||
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||
number of images that should be generated per prompt
|
||||
device: (`torch.device`, *optional*):
|
||||
torch device to place the resulting embeddings on
|
||||
prompt_embeds (`torch.FloatTensor`, *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.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`):
|
||||
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.
|
||||
"""
|
||||
embeds_initially_provided = prompt_embeds is not None and negative_prompt_embeds is not None
|
||||
|
||||
if device is None:
|
||||
device = self.text_encoder.device or self._execution_device
|
||||
|
||||
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]
|
||||
|
||||
# See Section 3.1. of the paper.
|
||||
max_length = 300
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt = self._text_preprocessing(prompt, clean_caption=clean_caption)
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_length,
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
untruncated_ids = self.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
|
||||
):
|
||||
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, max_length - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because the model can only handle sequences up to"
|
||||
f" {max_length} tokens: {removed_text}"
|
||||
)
|
||||
|
||||
attention_mask = text_inputs.attention_mask.to(device)
|
||||
prompt_embeds_attention_mask = attention_mask
|
||||
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=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
|
||||
elif self.transformer is not None:
|
||||
dtype = self.transformer.dtype
|
||||
else:
|
||||
dtype = None
|
||||
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
bs_embed, seq_len, _ = prompt_embeds.shape
|
||||
# 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)
|
||||
|
||||
# get unconditional embeddings for classifier free guidance
|
||||
if do_classifier_free_guidance and negative_prompt_embeds is None:
|
||||
uncond_tokens = [negative_prompt] * batch_size
|
||||
uncond_tokens = self._text_preprocessing(uncond_tokens, clean_caption=clean_caption)
|
||||
max_length = prompt_embeds.shape[1]
|
||||
uncond_input = self.tokenizer(
|
||||
uncond_tokens,
|
||||
padding="max_length",
|
||||
max_length=max_length,
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
attention_mask = uncond_input.attention_mask.to(device)
|
||||
|
||||
negative_prompt_embeds = self.text_encoder(
|
||||
uncond_input.input_ids.to(device),
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
negative_prompt_embeds = negative_prompt_embeds[0]
|
||||
|
||||
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)
|
||||
|
||||
# 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
|
||||
else:
|
||||
negative_prompt_embeds = None
|
||||
|
||||
# print(prompt_embeds.shape) # 1 120 4096
|
||||
# print(negative_prompt_embeds.shape) # 1 120 4096
|
||||
|
||||
# Perform additional masking.
|
||||
if mask_feature and not embeds_initially_provided:
|
||||
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
|
||||
)
|
||||
|
||||
# import torch.nn.functional as F
|
||||
|
||||
# padding = (0, 0, 0, 113) # (左, 右, 下, 上)
|
||||
# masked_prompt_embeds_ = F.pad(masked_prompt_embeds, padding, "constant", 0)
|
||||
# masked_negative_prompt_embeds_ = F.pad(masked_negative_prompt_embeds, padding, "constant", 0)
|
||||
|
||||
# print(masked_prompt_embeds == masked_prompt_embeds_[:, :masked_negative_prompt_embeds.shape[1], ...])
|
||||
|
||||
return masked_prompt_embeds, masked_negative_prompt_embeds
|
||||
# return masked_prompt_embeds_, masked_negative_prompt_embeds_
|
||||
|
||||
return prompt_embeds, negative_prompt_embeds
|
||||
|
||||
# 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,
|
||||
callback_steps,
|
||||
prompt_embeds=None,
|
||||
negative_prompt_embeds=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_steps is None) or (
|
||||
callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0)
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_steps` has to be a positive integer but is {callback_steps} of type"
|
||||
f" {type(callback_steps)}."
|
||||
)
|
||||
|
||||
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 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 is not None and negative_prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `negative_prompt_embeds`:"
|
||||
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
|
||||
)
|
||||
|
||||
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 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}."
|
||||
)
|
||||
|
||||
# Copied from diffusers.pipelines.deepfloyd_if.pipeline_if.IFPipeline._text_preprocessing
|
||||
def _text_preprocessing(self, text, clean_caption=False):
|
||||
if clean_caption and not is_bs4_available():
|
||||
logger.warn(BACKENDS_MAPPING["bs4"][-1].format("Setting `clean_caption=True`"))
|
||||
logger.warn("Setting `clean_caption` to False...")
|
||||
clean_caption = False
|
||||
|
||||
if clean_caption and not is_ftfy_available():
|
||||
logger.warn(BACKENDS_MAPPING["ftfy"][-1].format("Setting `clean_caption=True`"))
|
||||
logger.warn("Setting `clean_caption` to False...")
|
||||
clean_caption = False
|
||||
|
||||
if not isinstance(text, (tuple, list)):
|
||||
text = [text]
|
||||
|
||||
def process(text: str):
|
||||
if clean_caption:
|
||||
text = self._clean_caption(text)
|
||||
text = self._clean_caption(text)
|
||||
else:
|
||||
text = text.lower().strip()
|
||||
return text
|
||||
|
||||
return [process(t) for t in text]
|
||||
|
||||
# Copied from diffusers.pipelines.deepfloyd_if.pipeline_if.IFPipeline._clean_caption
|
||||
def _clean_caption(self, caption):
|
||||
caption = str(caption)
|
||||
caption = ul.unquote_plus(caption)
|
||||
caption = caption.strip().lower()
|
||||
caption = re.sub("<person>", "person", caption)
|
||||
# urls:
|
||||
caption = re.sub(
|
||||
r"\b((?:https?:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))",
|
||||
# noqa
|
||||
"",
|
||||
caption,
|
||||
) # regex for urls
|
||||
caption = re.sub(
|
||||
r"\b((?:www:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))",
|
||||
# noqa
|
||||
"",
|
||||
caption,
|
||||
) # regex for urls
|
||||
# html:
|
||||
caption = BeautifulSoup(caption, features="html.parser").text
|
||||
|
||||
# @<nickname>
|
||||
caption = re.sub(r"@[\w\d]+\b", "", caption)
|
||||
|
||||
# 31C0—31EF CJK Strokes
|
||||
# 31F0—31FF Katakana Phonetic Extensions
|
||||
# 3200—32FF Enclosed CJK Letters and Months
|
||||
# 3300—33FF CJK Compatibility
|
||||
# 3400—4DBF CJK Unified Ideographs Extension A
|
||||
# 4DC0—4DFF Yijing Hexagram Symbols
|
||||
# 4E00—9FFF CJK Unified Ideographs
|
||||
caption = re.sub(r"[\u31c0-\u31ef]+", "", caption)
|
||||
caption = re.sub(r"[\u31f0-\u31ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3200-\u32ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3300-\u33ff]+", "", caption)
|
||||
caption = re.sub(r"[\u3400-\u4dbf]+", "", caption)
|
||||
caption = re.sub(r"[\u4dc0-\u4dff]+", "", caption)
|
||||
caption = re.sub(r"[\u4e00-\u9fff]+", "", caption)
|
||||
#######################################################
|
||||
|
||||
# все виды тире / all types of dash --> "-"
|
||||
caption = re.sub(
|
||||
r"[\u002D\u058A\u05BE\u1400\u1806\u2010-\u2015\u2E17\u2E1A\u2E3A\u2E3B\u2E40\u301C\u3030\u30A0\uFE31\uFE32\uFE58\uFE63\uFF0D]+",
|
||||
# noqa
|
||||
"-",
|
||||
caption,
|
||||
)
|
||||
|
||||
# кавычки к одному стандарту
|
||||
caption = re.sub(r"[`´«»“”¨]", '"', caption)
|
||||
caption = re.sub(r"[‘’]", "'", caption)
|
||||
|
||||
# "
|
||||
caption = re.sub(r""?", "", caption)
|
||||
# &
|
||||
caption = re.sub(r"&", "", caption)
|
||||
|
||||
# ip adresses:
|
||||
caption = re.sub(r"\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}", " ", caption)
|
||||
|
||||
# article ids:
|
||||
caption = re.sub(r"\d:\d\d\s+$", "", caption)
|
||||
|
||||
# \n
|
||||
caption = re.sub(r"\\n", " ", caption)
|
||||
|
||||
# "#123"
|
||||
caption = re.sub(r"#\d{1,3}\b", "", caption)
|
||||
# "#12345.."
|
||||
caption = re.sub(r"#\d{5,}\b", "", caption)
|
||||
# "123456.."
|
||||
caption = re.sub(r"\b\d{6,}\b", "", caption)
|
||||
# filenames:
|
||||
caption = re.sub(r"[\S]+\.(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)", "", caption)
|
||||
|
||||
#
|
||||
caption = re.sub(r"[\"\']{2,}", r'"', caption) # """AUSVERKAUFT"""
|
||||
caption = re.sub(r"[\.]{2,}", r" ", caption) # """AUSVERKAUFT"""
|
||||
|
||||
caption = re.sub(self.bad_punct_regex, r" ", caption) # ***AUSVERKAUFT***, #AUSVERKAUFT
|
||||
caption = re.sub(r"\s+\.\s+", r" ", caption) # " . "
|
||||
|
||||
# this-is-my-cute-cat / this_is_my_cute_cat
|
||||
regex2 = re.compile(r"(?:\-|\_)")
|
||||
if len(re.findall(regex2, caption)) > 3:
|
||||
caption = re.sub(regex2, " ", caption)
|
||||
|
||||
caption = ftfy.fix_text(caption)
|
||||
caption = html.unescape(html.unescape(caption))
|
||||
|
||||
caption = re.sub(r"\b[a-zA-Z]{1,3}\d{3,15}\b", "", caption) # jc6640
|
||||
caption = re.sub(r"\b[a-zA-Z]+\d+[a-zA-Z]+\b", "", caption) # jc6640vc
|
||||
caption = re.sub(r"\b\d+[a-zA-Z]+\d+\b", "", caption) # 6640vc231
|
||||
|
||||
caption = re.sub(r"(worldwide\s+)?(free\s+)?shipping", "", caption)
|
||||
caption = re.sub(r"(free\s)?download(\sfree)?", "", caption)
|
||||
caption = re.sub(r"\bclick\b\s(?:for|on)\s\w+", "", caption)
|
||||
caption = re.sub(r"\b(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)(\simage[s]?)?", "", caption)
|
||||
caption = re.sub(r"\bpage\s+\d+\b", "", caption)
|
||||
|
||||
caption = re.sub(r"\b\d*[a-zA-Z]+\d+[a-zA-Z]+\d+[a-zA-Z\d]*\b", r" ", caption) # j2d1a2a...
|
||||
|
||||
caption = re.sub(r"\b\d+\.?\d*[xх×]\d+\.?\d*\b", "", caption)
|
||||
|
||||
caption = re.sub(r"\b\s+\:\s+", r": ", caption)
|
||||
caption = re.sub(r"(\D[,\./])\b", r"\1 ", caption)
|
||||
caption = re.sub(r"\s+", " ", caption)
|
||||
|
||||
caption.strip()
|
||||
|
||||
caption = re.sub(r"^[\"\']([\w\W]+)[\"\']$", r"\1", caption)
|
||||
caption = re.sub(r"^[\'\_,\-\:;]", r"", caption)
|
||||
caption = re.sub(r"[\'\_,\-\:\-\+]$", r"", caption)
|
||||
caption = re.sub(r"^\.\S+$", "", caption)
|
||||
|
||||
return caption.strip()
|
||||
|
||||
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_latents
|
||||
def prepare_latents(
|
||||
self, batch_size, num_channels_latents, num_frames, height, width, dtype, device, generator, latents=None
|
||||
):
|
||||
shape = (
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
(math.ceil((int(num_frames) - 1) / self.vae.vae_scale_factor[0]) + 1)
|
||||
if int(num_frames) % 2 == 1
|
||||
else math.ceil(int(num_frames) / self.vae.vae_scale_factor[0]),
|
||||
math.ceil(int(height) / self.vae.vae_scale_factor[1]),
|
||||
math.ceil(int(width) / self.vae.vae_scale_factor[2]),
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
negative_prompt: str = "",
|
||||
num_inference_steps: int = 20,
|
||||
timesteps: List[int] = None,
|
||||
guidance_scale: float = 4.5,
|
||||
num_images_per_prompt: Optional[int] = 1,
|
||||
num_frames: Optional[int] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
eta: float = 0.0,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
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,
|
||||
enable_temporal_attentions: bool = True,
|
||||
verbose: bool = False,
|
||||
) -> Union[VideoPipelineOutput, Tuple]:
|
||||
"""
|
||||
Function invoked when calling the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
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`).
|
||||
num_inference_steps (`int`, *optional*, defaults to 100):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
timesteps (`List[int]`, *optional*):
|
||||
Custom timesteps to use for the denoising process. If not defined, equal spaced `num_inference_steps`
|
||||
timesteps are used. Must be in descending order.
|
||||
guidance_scale (`float`, *optional*, defaults to 7.0):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
|
||||
`guidance_scale` is defined as `w` of equation 2. of [Imagen
|
||||
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
|
||||
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
|
||||
usually at the expense of lower image quality.
|
||||
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
height (`int`, *optional*, defaults to self.unet.config.sample_size):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, *optional*, defaults to self.unet.config.sample_size):
|
||||
The width in pixels of the generated image.
|
||||
eta (`float`, *optional*, defaults to 0.0):
|
||||
Corresponds to parameter eta (η) in the DDIM paper: https://arxiv.org/abs/2010.02502. Only applies to
|
||||
[`schedulers.DDIMScheduler`], will be ignored for others.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
|
||||
to make generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor will ge generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.FloatTensor`, *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.FloatTensor`, *optional*):
|
||||
Pre-generated negative text embeddings. For PixArt-Alpha this negative prompt should be "". If not
|
||||
provided, negative_prompt_embeds will be generated from `negative_prompt` input argument.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generate image. Choose between
|
||||
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~pipelines.stable_diffusion.IFPipelineOutput`] instead of a plain tuple.
|
||||
callback (`Callable`, *optional*):
|
||||
A function that will be called every `callback_steps` steps during inference. The function will be
|
||||
called with the following arguments: `callback(step: int, timestep: int, latents: torch.FloatTensor)`.
|
||||
callback_steps (`int`, *optional*, defaults to 1):
|
||||
The frequency at which the `callback` function will be called. If not specified, the callback will be
|
||||
called at every step.
|
||||
clean_caption (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to clean the caption before creating embeddings. Requires `beautifulsoup4` and `ftfy` to
|
||||
be installed. If the dependencies are not installed, the embeddings will be created from the raw
|
||||
prompt.
|
||||
mask_feature (`bool` defaults to `True`): If set to `True`, the text embeddings will be masked.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~pipelines.ImagePipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`, [`~pipelines.ImagePipelineOutput`] is returned, otherwise a `tuple` is
|
||||
returned where the first element is a list with the generated images
|
||||
"""
|
||||
# 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
|
||||
self.check_inputs(prompt, height, width, negative_prompt, callback_steps, prompt_embeds, negative_prompt_embeds)
|
||||
|
||||
# 2. Default height and width to transformer
|
||||
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.text_encoder.device or self._execution_device
|
||||
|
||||
# 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.
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
|
||||
# 3. Encode input prompt
|
||||
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
|
||||
prompt,
|
||||
do_classifier_free_guidance,
|
||||
negative_prompt=negative_prompt,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
device=device,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
clean_caption=clean_caption,
|
||||
mask_feature=mask_feature,
|
||||
)
|
||||
if do_classifier_free_guidance:
|
||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 5. Prepare latents.
|
||||
latent_channels = self.transformer.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_images_per_prompt,
|
||||
latent_channels,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
# 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)
|
||||
|
||||
# 6.1 Prepare micro-conditions.
|
||||
added_cond_kwargs = {"resolution": None, "aspect_ratio": None}
|
||||
# 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)
|
||||
# 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)
|
||||
|
||||
if get_diffusion_skip() and get_diffusion_skip_timestep() is not None:
|
||||
diffusion_skip_timestep = get_diffusion_skip_timestep()
|
||||
|
||||
# warmup_timesteps = timesteps[:num_warmup_steps]
|
||||
# after_warmup_timesteps = skip_diffusion_timestep(timesteps[num_warmup_steps:], diffusion_skip_timestep)
|
||||
# timesteps = torch.cat((warmup_timesteps, after_warmup_timesteps))
|
||||
|
||||
timesteps = skip_diffusion_timestep(timesteps, diffusion_skip_timestep)
|
||||
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
orignal_timesteps = self.scheduler.timesteps
|
||||
|
||||
if verbose and dist.get_rank() == 0:
|
||||
print("============================")
|
||||
print(f"orignal sample timesteps: {orignal_timesteps}")
|
||||
print(f"orignal diffusion steps: {len(orignal_timesteps)}")
|
||||
print("============================")
|
||||
print(f"skip diffusion steps: {get_diffusion_skip_timestep()}")
|
||||
print(f"sample timesteps: {timesteps}")
|
||||
print(f"num_inference_steps: {len(timesteps)}")
|
||||
print("============================")
|
||||
|
||||
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)
|
||||
|
||||
current_timestep = t
|
||||
if not torch.is_tensor(current_timestep):
|
||||
# TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can
|
||||
# This would be a good case for the `match` statement (Python 3.10+)
|
||||
is_mps = latent_model_input.device.type == "mps"
|
||||
if isinstance(current_timestep, float):
|
||||
dtype = torch.float32 if is_mps else torch.float64
|
||||
else:
|
||||
dtype = torch.int32 if is_mps else torch.int64
|
||||
current_timestep = torch.tensor([current_timestep], dtype=dtype, device=latent_model_input.device)
|
||||
elif len(current_timestep.shape) == 0:
|
||||
current_timestep = current_timestep[None].to(latent_model_input.device)
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
current_timestep = current_timestep.expand(latent_model_input.shape[0])
|
||||
|
||||
if prompt_embeds.ndim == 3:
|
||||
prompt_embeds = prompt_embeds.unsqueeze(1) # b l d -> b 1 l d
|
||||
# if prompt_attention_mask.ndim == 2:
|
||||
# prompt_attention_mask = prompt_attention_mask.unsqueeze(1) # b l -> b 1 l
|
||||
# predict noise model_output
|
||||
noise_pred = self.transformer(
|
||||
latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep=current_timestep,
|
||||
added_cond_kwargs=added_cond_kwargs,
|
||||
enable_temporal_attentions=enable_temporal_attentions,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# perform guidance
|
||||
if 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)
|
||||
|
||||
# learned sigma
|
||||
if self.transformer.config.out_channels // 2 == latent_channels:
|
||||
noise_pred = noise_pred.chunk(2, dim=1)[0]
|
||||
else:
|
||||
noise_pred = noise_pred
|
||||
|
||||
# compute previous image: x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
|
||||
# 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()
|
||||
if callback is not None and i % callback_steps == 0:
|
||||
step_idx = i // getattr(self.scheduler, "order", 1)
|
||||
callback(step_idx, t, latents)
|
||||
|
||||
if not output_type == "latents":
|
||||
video = self.decode_latents(latents)
|
||||
video = video[:, :num_frames, :height, :width]
|
||||
else:
|
||||
video = latents
|
||||
return VideoPipelineOutput(video=video)
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (video,)
|
||||
|
||||
return VideoPipelineOutput(video=video)
|
||||
|
||||
def decode_latents(self, latents):
|
||||
video = self.vae.decode(latents) # b t c h w
|
||||
# b t c h w -> b t h w c
|
||||
video = ((video / 2.0 + 0.5).clamp(0, 1) * 255).to(dtype=torch.uint8).cpu().permute(0, 1, 3, 4, 2).contiguous()
|
||||
return video
|
||||
Executable
Executable
+217
@@ -0,0 +1,217 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Iterable, List, Optional, Sequence, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.checkpoint
|
||||
|
||||
from opendit.modules.layers import LlamaRMSNorm
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_heads: int = 8,
|
||||
qkv_bias: bool = False,
|
||||
qk_norm: bool = False,
|
||||
attn_drop: float = 0.0,
|
||||
proj_drop: float = 0.0,
|
||||
norm_layer: nn.Module = LlamaRMSNorm,
|
||||
enable_flashattn: bool = False,
|
||||
rope=None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
assert dim % num_heads == 0, "dim should be divisible by num_heads"
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.scale = self.head_dim**-0.5
|
||||
self.enable_flashattn = enable_flashattn
|
||||
|
||||
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
self.q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
|
||||
self.k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
|
||||
self.rope = False
|
||||
if rope is not None:
|
||||
self.rope = True
|
||||
self.rotary_emb = rope
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
B, N, C = x.shape
|
||||
|
||||
qkv = self.qkv(x)
|
||||
qkv = qkv.view(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 1, 3, 4)
|
||||
q, k, v = qkv.unbind(0)
|
||||
if self.rope:
|
||||
q = self.rotary_emb(q)
|
||||
k = self.rotary_emb(k)
|
||||
q, k = self.q_norm(q), self.k_norm(k)
|
||||
|
||||
if self.enable_flashattn:
|
||||
from flash_attn import flash_attn_func
|
||||
|
||||
x = flash_attn_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=self.attn_drop.p if self.training else 0.0,
|
||||
softmax_scale=self.scale,
|
||||
)
|
||||
else:
|
||||
q, k, v = map(lambda t: t.permute(0, 2, 1, 3), (q, k, v))
|
||||
x = F.scaled_dot_product_attention(
|
||||
q, k, v, scale=self.scale, dropout_p=self.attn_drop.p if self.training else 0.0
|
||||
)
|
||||
|
||||
x_output_shape = (B, N, C)
|
||||
if not self.enable_flashattn:
|
||||
x = x.transpose(1, 2)
|
||||
x = x.reshape(x_output_shape)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class MultiHeadCrossAttention(nn.Module):
|
||||
def __init__(self, d_model, num_heads, attn_drop=0.0, proj_drop=0.0, enable_flashattn=False):
|
||||
super(MultiHeadCrossAttention, self).__init__()
|
||||
assert d_model % num_heads == 0, "d_model must be divisible by num_heads"
|
||||
|
||||
self.d_model = d_model
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = d_model // num_heads
|
||||
self.enable_flashattn = enable_flashattn
|
||||
|
||||
self.q_linear = nn.Linear(d_model, d_model)
|
||||
self.kv_linear = nn.Linear(d_model, d_model * 2)
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(d_model, d_model)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
self.last_out = None
|
||||
self.count = 0
|
||||
|
||||
def forward(self, x, cond, mask=None, timestep=None):
|
||||
# query/value: img tokens; key: condition; mask: if padding tokens
|
||||
B, N, C = x.shape
|
||||
|
||||
q = self.q_linear(x).view(1, -1, self.num_heads, self.head_dim)
|
||||
kv = self.kv_linear(cond).view(1, -1, 2, self.num_heads, self.head_dim)
|
||||
k, v = kv.unbind(2)
|
||||
x = self.flash_attn_impl(q, k, v, mask, B, N, C)
|
||||
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
return x
|
||||
|
||||
def flash_attn_impl(self, q, k, v, mask, B, N, C):
|
||||
from flash_attn import flash_attn_varlen_func
|
||||
|
||||
q_seqinfo = _SeqLenInfo.from_seqlens([N] * B)
|
||||
k_seqinfo = _SeqLenInfo.from_seqlens(mask)
|
||||
|
||||
x = flash_attn_varlen_func(
|
||||
q.view(-1, self.num_heads, self.head_dim),
|
||||
k.view(-1, self.num_heads, self.head_dim),
|
||||
v.view(-1, self.num_heads, self.head_dim),
|
||||
cu_seqlens_q=q_seqinfo.seqstart.cuda(),
|
||||
cu_seqlens_k=k_seqinfo.seqstart.cuda(),
|
||||
max_seqlen_q=q_seqinfo.max_seqlen,
|
||||
max_seqlen_k=k_seqinfo.max_seqlen,
|
||||
dropout_p=self.attn_drop.p if self.training else 0.0,
|
||||
)
|
||||
x = x.view(B, N, C)
|
||||
return x
|
||||
|
||||
def torch_impl(self, q, k, v, mask, B, N, C):
|
||||
q = q.view(B, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
||||
k = k.view(B, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
||||
v = v.view(B, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
||||
|
||||
attn_mask = torch.zeros(B, N, k.shape[2], dtype=torch.float32, device=q.device)
|
||||
for i, m in enumerate(mask):
|
||||
attn_mask[i, :, m:] = -1e8
|
||||
|
||||
scale = 1 / q.shape[-1] ** 0.5
|
||||
q = q * scale
|
||||
attn = q @ k.transpose(-2, -1)
|
||||
attn = attn.to(torch.float32)
|
||||
if mask is not None:
|
||||
attn = attn + attn_mask.unsqueeze(1)
|
||||
attn = attn.softmax(-1)
|
||||
attn = attn.to(v.dtype)
|
||||
out = attn @ v
|
||||
|
||||
x = out.transpose(1, 2).contiguous().view(B, N, C)
|
||||
return x
|
||||
|
||||
|
||||
@dataclass
|
||||
class _SeqLenInfo:
|
||||
"""
|
||||
copied from xformers
|
||||
|
||||
(Internal) Represents the division of a dimension into blocks.
|
||||
For example, to represents a dimension of length 7 divided into
|
||||
three blocks of lengths 2, 3 and 2, use `from_seqlength([2, 3, 2])`.
|
||||
The members will be:
|
||||
max_seqlen: 3
|
||||
min_seqlen: 2
|
||||
seqstart_py: [0, 2, 5, 7]
|
||||
seqstart: torch.IntTensor([0, 2, 5, 7])
|
||||
"""
|
||||
|
||||
seqstart: torch.Tensor
|
||||
max_seqlen: int
|
||||
min_seqlen: int
|
||||
seqstart_py: List[int]
|
||||
|
||||
def to(self, device: torch.device) -> None:
|
||||
self.seqstart = self.seqstart.to(device, non_blocking=True)
|
||||
|
||||
def intervals(self) -> Iterable[Tuple[int, int]]:
|
||||
yield from zip(self.seqstart_py, self.seqstart_py[1:])
|
||||
|
||||
@classmethod
|
||||
def from_seqlens(cls, seqlens: Iterable[int]) -> "_SeqLenInfo":
|
||||
"""
|
||||
Input tensors are assumed to be in shape [B, M, *]
|
||||
"""
|
||||
assert not isinstance(seqlens, torch.Tensor)
|
||||
seqstart_py = [0]
|
||||
max_seqlen = -1
|
||||
min_seqlen = -1
|
||||
for seqlen in seqlens:
|
||||
min_seqlen = min(min_seqlen, seqlen) if min_seqlen != -1 else seqlen
|
||||
max_seqlen = max(max_seqlen, seqlen)
|
||||
seqstart_py.append(seqstart_py[len(seqstart_py) - 1] + seqlen)
|
||||
seqstart = torch.tensor(seqstart_py, dtype=torch.int32)
|
||||
return cls(
|
||||
max_seqlen=max_seqlen,
|
||||
min_seqlen=min_seqlen,
|
||||
seqstart=seqstart,
|
||||
seqstart_py=seqstart_py,
|
||||
)
|
||||
|
||||
def split(self, x: torch.Tensor, batch_sizes: Optional[Sequence[int]] = None) -> List[torch.Tensor]:
|
||||
if self.seqstart_py[-1] != x.shape[1] or x.shape[0] != 1:
|
||||
raise ValueError(
|
||||
f"Invalid `torch.Tensor` of shape {x.shape}, expected format "
|
||||
f"(B, M, *) with B=1 and M={self.seqstart_py[-1]}\n"
|
||||
f" seqstart: {self.seqstart_py}"
|
||||
)
|
||||
if batch_sizes is None:
|
||||
batch_sizes = [1] * (len(self.seqstart_py) - 1)
|
||||
split_chunks = []
|
||||
it = 0
|
||||
for batch_size in batch_sizes:
|
||||
split_chunks.append(self.seqstart_py[it + batch_size] - self.seqstart_py[it])
|
||||
it += batch_size
|
||||
return [
|
||||
tensor.reshape([bs, -1, *tensor.shape[2:]]) for bs, tensor in zip(batch_sizes, x.split(split_chunks, dim=1))
|
||||
]
|
||||
Executable
+145
@@ -0,0 +1,145 @@
|
||||
# Modified from Meta DiT
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# DiT: https://github.com/facebookresearch/DiT/tree/main
|
||||
# GLIDE: https://github.com/openai/glide-text2im
|
||||
# MAE: https://github.com/facebookresearch/mae/blob/main/models_mae.py
|
||||
# --------------------------------------------------------
|
||||
|
||||
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__()
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
||||
)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
|
||||
@staticmethod
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param t: 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, D) Tensor of positional embeddings.
|
||||
"""
|
||||
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
||||
half = dim // 2
|
||||
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to(
|
||||
device=t.device
|
||||
)
|
||||
args = t[:, 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)
|
||||
return embedding
|
||||
|
||||
def forward(self, t):
|
||||
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
|
||||
|
||||
class LabelEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
||||
"""
|
||||
|
||||
def __init__(self, num_classes, hidden_size, dropout_prob):
|
||||
super().__init__()
|
||||
use_cfg_embedding = dropout_prob > 0
|
||||
self.embedding_table = nn.Embedding(num_classes + use_cfg_embedding, hidden_size)
|
||||
self.num_classes = num_classes
|
||||
self.dropout_prob = dropout_prob
|
||||
|
||||
def token_drop(self, labels, force_drop_ids=None):
|
||||
"""
|
||||
Drops labels to enable classifier-free guidance.
|
||||
"""
|
||||
if force_drop_ids is None:
|
||||
drop_ids = torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob
|
||||
else:
|
||||
drop_ids = force_drop_ids == 1
|
||||
labels = torch.where(drop_ids, self.num_classes, labels)
|
||||
return labels
|
||||
|
||||
def forward(self, labels, train, force_drop_ids=None):
|
||||
use_dropout = self.dropout_prob > 0
|
||||
if (train and use_dropout) or (force_drop_ids is not None):
|
||||
labels = self.token_drop(labels, force_drop_ids)
|
||||
embeddings = self.embedding_table(labels)
|
||||
return embeddings
|
||||
|
||||
|
||||
#################################################################################
|
||||
# Sine/Cosine Positional Embedding Functions #
|
||||
#################################################################################
|
||||
# https://github.com/facebookresearch/mae/blob/main/util/pos_embed.py
|
||||
|
||||
|
||||
def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0):
|
||||
"""
|
||||
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 and extra_tokens > 0:
|
||||
pos_embed = np.concatenate([np.zeros([extra_tokens, 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.float64)
|
||||
omega /= embed_dim / 2.0
|
||||
omega = 1.0 / 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
|
||||
Executable
+80
@@ -0,0 +1,80 @@
|
||||
# Modified from Meta DiT
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# DiT: https://github.com/facebookresearch/DiT/tree/main
|
||||
# GLIDE: https://github.com/openai/glide-text2im
|
||||
# MAE: https://github.com/facebookresearch/mae/blob/main/models_mae.py
|
||||
# --------------------------------------------------------
|
||||
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.utils.checkpoint
|
||||
|
||||
|
||||
def get_layernorm(hidden_size: torch.Tensor, eps: float, affine: bool, use_kernel: bool):
|
||||
if use_kernel:
|
||||
try:
|
||||
from apex.normalization import FusedLayerNorm
|
||||
|
||||
return FusedLayerNorm(hidden_size, elementwise_affine=affine, eps=eps)
|
||||
except ImportError:
|
||||
raise RuntimeError("FusedLayerNorm not available. Please install apex.")
|
||||
else:
|
||||
return nn.LayerNorm(hidden_size, eps, elementwise_affine=affine)
|
||||
|
||||
|
||||
def modulate(norm_func, x, shift, scale, use_kernel=False):
|
||||
# Suppose x is (N, T, D), shift is (N, D), scale is (N, D)
|
||||
dtype = x.dtype
|
||||
x = norm_func(x.to(torch.float32)).to(dtype)
|
||||
if use_kernel:
|
||||
try:
|
||||
from opendit.kernels.fused_modulate import fused_modulate
|
||||
|
||||
x = fused_modulate(x, scale, shift)
|
||||
except ImportError:
|
||||
raise RuntimeError("FusedModulate kernel not available. Please install triton.")
|
||||
else:
|
||||
x = x * (scale.unsqueeze(1) + 1) + shift.unsqueeze(1)
|
||||
x = x.to(dtype)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class FinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of DiT.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, patch_size, out_channels):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size, patch_size * out_channels, bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
|
||||
def forward(self, x, c):
|
||||
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
|
||||
x = modulate(self.norm_final, x, shift, scale)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class LlamaRMSNorm(nn.Module):
|
||||
def __init__(self, hidden_size, eps=1e-6):
|
||||
"""
|
||||
LlamaRMSNorm is equivalent to T5LayerNorm
|
||||
"""
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(hidden_size))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
def forward(self, hidden_states):
|
||||
input_dtype = hidden_states.dtype
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
variance = hidden_states.pow(2).mean(-1, keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
|
||||
return self.weight * hidden_states.to(input_dtype)
|
||||
Executable
+134
@@ -0,0 +1,134 @@
|
||||
import functools
|
||||
import json
|
||||
import logging
|
||||
import operator
|
||||
import os
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from colossalai.booster import Booster
|
||||
from colossalai.cluster import DistCoordinator
|
||||
from torch.optim import Optimizer
|
||||
from torch.optim.lr_scheduler import _LRScheduler
|
||||
|
||||
from opendit.core.comm import model_sharding
|
||||
|
||||
|
||||
def load_json(file_path: str):
|
||||
with open(file_path, "r") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def save_json(data, file_path: str):
|
||||
with open(file_path, "w") as f:
|
||||
json.dump(data, f, indent=4)
|
||||
|
||||
|
||||
def remove_padding(tensor: torch.Tensor, original_shape: Tuple) -> torch.Tensor:
|
||||
return tensor[: functools.reduce(operator.mul, original_shape)]
|
||||
|
||||
|
||||
def model_gathering(model: torch.nn.Module, model_shape_dict: dict):
|
||||
global_rank = dist.get_rank()
|
||||
global_size = dist.get_world_size()
|
||||
for name, param in model.named_parameters():
|
||||
all_params = [torch.empty_like(param.data) for _ in range(global_size)]
|
||||
dist.all_gather(all_params, param.data, group=dist.group.WORLD)
|
||||
if global_rank == 0:
|
||||
all_params = torch.cat(all_params)
|
||||
param.data = remove_padding(all_params, model_shape_dict[name]).view(model_shape_dict[name])
|
||||
dist.barrier()
|
||||
|
||||
|
||||
def record_model_param_shape(model: torch.nn.Module) -> dict:
|
||||
param_shape = {}
|
||||
for name, param in model.named_parameters():
|
||||
param_shape[name] = param.shape
|
||||
return param_shape
|
||||
|
||||
|
||||
def save(
|
||||
booster: Booster,
|
||||
model: nn.Module,
|
||||
ema: nn.Module,
|
||||
optimizer: Optimizer,
|
||||
lr_scheduler: _LRScheduler,
|
||||
epoch: int,
|
||||
step: int,
|
||||
global_step: int,
|
||||
batch_size: int,
|
||||
coordinator: DistCoordinator,
|
||||
save_dir: str,
|
||||
shape_dict: dict,
|
||||
shard_ema: bool = False,
|
||||
):
|
||||
torch.cuda.empty_cache()
|
||||
global_rank = dist.get_rank()
|
||||
save_dir = os.path.join(save_dir, f"epoch{epoch}-global_step{global_step}")
|
||||
os.makedirs(os.path.join(save_dir, "model"), exist_ok=True)
|
||||
booster.save_model(model, os.path.join(save_dir, "model"), shard=True)
|
||||
|
||||
# Gather the sharded ema model before saving
|
||||
if shard_ema:
|
||||
model_gathering(ema, shape_dict)
|
||||
|
||||
# ema is not boosted, so we don't need to use booster.save_model
|
||||
if global_rank == 0:
|
||||
torch.save(ema.state_dict(), os.path.join(save_dir, "ema.pt"))
|
||||
# Shard ema model when using zero2 plugin
|
||||
if shard_ema:
|
||||
model_sharding(ema)
|
||||
if optimizer is not None:
|
||||
booster.save_optimizer(optimizer, os.path.join(save_dir, "optimizer"), shard=True, size_per_shard=4096)
|
||||
if lr_scheduler is not None:
|
||||
booster.save_lr_scheduler(lr_scheduler, os.path.join(save_dir, "lr_scheduler"))
|
||||
running_states = {
|
||||
"epoch": epoch,
|
||||
"step": step,
|
||||
"global_step": global_step,
|
||||
"sample_start_index": step * batch_size,
|
||||
}
|
||||
if coordinator.is_master():
|
||||
save_json(running_states, os.path.join(save_dir, "running_states.json"))
|
||||
dist.barrier()
|
||||
|
||||
|
||||
def load(
|
||||
booster: Booster,
|
||||
model: nn.Module,
|
||||
ema: nn.Module,
|
||||
optimizer: Optimizer,
|
||||
lr_scheduler: _LRScheduler,
|
||||
load_dir: str,
|
||||
) -> Tuple[int, int, int]:
|
||||
booster.load_model(model, os.path.join(load_dir, "model"))
|
||||
# ema is not boosted, so we don't use booster.load_model
|
||||
ema.load_state_dict(torch.load(os.path.join(load_dir, "ema.pt"), map_location=torch.device("cpu")))
|
||||
if optimizer is not None:
|
||||
booster.load_optimizer(optimizer, os.path.join(load_dir, "optimizer"))
|
||||
if lr_scheduler is not None:
|
||||
booster.load_lr_scheduler(lr_scheduler, os.path.join(load_dir, "lr_scheduler"))
|
||||
running_states = load_json(os.path.join(load_dir, "running_states.json"))
|
||||
dist.barrier()
|
||||
torch.cuda.empty_cache()
|
||||
return running_states["epoch"], running_states["step"], running_states["sample_start_index"]
|
||||
|
||||
|
||||
def create_logger(logging_dir):
|
||||
"""
|
||||
Create a logger that writes to a log file and stdout.
|
||||
"""
|
||||
if dist.get_rank() == 0: # real logger
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="[\033[34m%(asctime)s\033[0m] %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
handlers=[logging.StreamHandler(), logging.FileHandler(f"{logging_dir}/log.txt")],
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
else: # dummy logger (does nothing)
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.addHandler(logging.NullHandler())
|
||||
return logger
|
||||
Executable
+7
@@ -0,0 +1,7 @@
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
# Print debug information on selected rank
|
||||
def print_rank(var_name, var_value, rank=0):
|
||||
if dist.get_rank() == rank:
|
||||
print(f"[Rank {rank}] {var_name}: {var_value}")
|
||||
Executable
+79
@@ -0,0 +1,79 @@
|
||||
# Copyright (c) Meta Platforms, Inc. and affiliates.
|
||||
# All rights reserved.
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
"""
|
||||
Functions for downloading pre-trained DiT models
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
from torchvision.datasets.utils import download_url
|
||||
|
||||
pretrained_models = {"DiT-XL-2-512x512.pt", "DiT-XL-2-256x256.pt"}
|
||||
|
||||
|
||||
def find_model(model_name):
|
||||
"""
|
||||
Finds a pre-trained DiT model, downloading it if necessary. Alternatively, loads a model from a local path.
|
||||
"""
|
||||
if model_name in pretrained_models: # Find/download our pre-trained DiT checkpoints
|
||||
return download_model(model_name)
|
||||
else: # Load a custom DiT checkpoint:
|
||||
if not os.path.isfile(model_name):
|
||||
# if the model_name is a directory, then we assume we should load it in the Hugging Face manner
|
||||
# i.e. the model weights are sharded into multiple files and there is an index.json file
|
||||
# walk through the files in the directory and find the index.json file
|
||||
index_file = [os.path.join(model_name, f) for f in os.listdir(model_name) if "index.json" in f]
|
||||
assert len(index_file) == 1, f"Could not find index.json in {model_name}"
|
||||
|
||||
# process index json
|
||||
with open(index_file[0], "r") as f:
|
||||
index_data = json.load(f)
|
||||
|
||||
bin_to_weight_mapping = dict()
|
||||
for k, v in index_data["weight_map"].items():
|
||||
if v in bin_to_weight_mapping:
|
||||
bin_to_weight_mapping[v].append(k)
|
||||
else:
|
||||
bin_to_weight_mapping[v] = [k]
|
||||
|
||||
# make state dict
|
||||
state_dict = dict()
|
||||
for bin_name, weight_list in bin_to_weight_mapping.items():
|
||||
bin_path = os.path.join(model_name, bin_name)
|
||||
bin_state_dict = torch.load(bin_path, map_location=lambda storage, loc: storage)
|
||||
for weight in weight_list:
|
||||
state_dict[weight] = bin_state_dict[weight]
|
||||
return state_dict
|
||||
else:
|
||||
# if it is a file, we just load it directly in the typical PyTorch manner
|
||||
assert os.path.exists(model_name), f"Could not find DiT checkpoint at {model_name}"
|
||||
checkpoint = torch.load(model_name, map_location=lambda storage, loc: storage)
|
||||
if "ema" in checkpoint: # supports checkpoints from train.py
|
||||
checkpoint = checkpoint["ema"]
|
||||
return checkpoint
|
||||
|
||||
|
||||
def download_model(model_name):
|
||||
"""
|
||||
Downloads a pre-trained DiT model from the web.
|
||||
"""
|
||||
assert model_name in pretrained_models
|
||||
local_path = f"pretrained_models/{model_name}"
|
||||
if not os.path.isfile(local_path):
|
||||
os.makedirs("pretrained_models", exist_ok=True)
|
||||
web_path = f"https://dl.fbaipublicfiles.com/DiT/models/{model_name}"
|
||||
download_url(web_path, "pretrained_models")
|
||||
model = torch.load(local_path, map_location=lambda storage, loc: storage)
|
||||
return model
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Download all DiT checkpoints
|
||||
for model in pretrained_models:
|
||||
download_model(model)
|
||||
print("Done.")
|
||||
Executable
+65
@@ -0,0 +1,65 @@
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from colossalai.zero.low_level.low_level_optim import LowLevelZeroOptimizer
|
||||
|
||||
|
||||
def get_model_numel(model: torch.nn.Module) -> int:
|
||||
return sum(p.numel() for p in model.parameters())
|
||||
|
||||
|
||||
def format_numel_str(numel: int) -> str:
|
||||
B = 1024**3
|
||||
M = 1024**2
|
||||
K = 1024
|
||||
if numel >= B:
|
||||
return f"{numel / B:.2f} B"
|
||||
elif numel >= M:
|
||||
return f"{numel / M:.2f} M"
|
||||
elif numel >= K:
|
||||
return f"{numel / K:.2f} K"
|
||||
else:
|
||||
return f"{numel}"
|
||||
|
||||
|
||||
def all_reduce_mean(tensor: torch.Tensor) -> torch.Tensor:
|
||||
dist.all_reduce(tensor=tensor, op=dist.ReduceOp.SUM)
|
||||
tensor.div_(dist.get_world_size())
|
||||
return tensor
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def update_ema(
|
||||
ema_model: torch.nn.Module, model: torch.nn.Module, optimizer=None, decay: float = 0.9999, sharded: bool = True
|
||||
) -> None:
|
||||
"""
|
||||
Step the EMA model towards the current model.
|
||||
"""
|
||||
ema_params = OrderedDict(ema_model.named_parameters())
|
||||
model_params = OrderedDict(model.named_parameters())
|
||||
|
||||
for name, param in model_params.items():
|
||||
if name == "pos_embed":
|
||||
continue
|
||||
if param.requires_grad == False:
|
||||
continue
|
||||
if not sharded:
|
||||
param_data = param.data
|
||||
ema_params[name].mul_(decay).add_(param_data, alpha=1 - decay)
|
||||
else:
|
||||
if param.data.dtype != torch.float32 and isinstance(optimizer, LowLevelZeroOptimizer):
|
||||
param_id = id(param)
|
||||
master_param = optimizer._param_store.working_to_master_param[param_id]
|
||||
param_data = master_param.data
|
||||
else:
|
||||
param_data = param.data
|
||||
ema_params[name].mul_(decay).add_(param_data, alpha=1 - decay)
|
||||
|
||||
|
||||
def requires_grad(model: torch.nn.Module, flag: bool = True) -> None:
|
||||
"""
|
||||
Set requires_grad flag for all parameters in a model.
|
||||
"""
|
||||
for p in model.parameters():
|
||||
p.requires_grad = flag
|
||||
Executable
+86
@@ -0,0 +1,86 @@
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from omegaconf import DictConfig, ListConfig, OmegaConf
|
||||
|
||||
|
||||
def requires_grad(model: torch.nn.Module, flag: bool = True) -> None:
|
||||
"""
|
||||
Set requires_grad flag for all parameters in a model.
|
||||
"""
|
||||
for p in model.parameters():
|
||||
p.requires_grad = flag
|
||||
|
||||
|
||||
def set_seed(seed):
|
||||
random.seed(seed)
|
||||
os.environ["PYTHONHASHSEED"] = str(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
|
||||
def str_to_dtype(x: str):
|
||||
if x == "fp32":
|
||||
return torch.float32
|
||||
elif x == "fp16":
|
||||
return torch.float16
|
||||
elif x == "bf16":
|
||||
return torch.bfloat16
|
||||
else:
|
||||
raise RuntimeError(f"Only fp32, fp16 and bf16 are supported, but got {x}")
|
||||
|
||||
|
||||
def merge_args(args1, args2):
|
||||
"""
|
||||
Merge two argparse Namespace objects.
|
||||
"""
|
||||
if args2 is None:
|
||||
return args1
|
||||
|
||||
for k in args2._content.keys():
|
||||
if k in args1.__dict__:
|
||||
v = getattr(args2, k)
|
||||
if isinstance(v, ListConfig) or isinstance(v, DictConfig):
|
||||
v = OmegaConf.to_object(v)
|
||||
setattr(args1, k, v)
|
||||
else:
|
||||
raise RuntimeError(f"Unknown argument {k}")
|
||||
|
||||
return args1
|
||||
|
||||
|
||||
def all_exists(paths):
|
||||
return all(os.path.exists(path) for path in paths)
|
||||
|
||||
|
||||
def get_logger():
|
||||
return logging.getLogger(__name__)
|
||||
|
||||
|
||||
def create_logger(logging_dir=None):
|
||||
"""
|
||||
Create a logger that writes to a log file and stdout.
|
||||
"""
|
||||
if dist.get_rank() == 0:
|
||||
additional_args = dict()
|
||||
if logging_dir is not None:
|
||||
additional_args["handlers"] = [
|
||||
logging.StreamHandler(),
|
||||
logging.FileHandler(f"{logging_dir}/log.txt"),
|
||||
]
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="[\033[34m%(asctime)s\033[0m] %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
**additional_args,
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
else: # dummy logger (does nothing)
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.addHandler(logging.NullHandler())
|
||||
return logger
|
||||
Executable
+7
@@ -0,0 +1,7 @@
|
||||
numpy
|
||||
timm
|
||||
accelerate
|
||||
diffusers
|
||||
transformers
|
||||
rotary_embedding_torch
|
||||
bs4
|
||||
Executable
+105
@@ -0,0 +1,105 @@
|
||||
# Modified from Meta DiT: https://github.com/facebookresearch/DiT
|
||||
|
||||
# Copyright (c) Meta Platforms, Inc. and affiliates.
|
||||
# All rights reserved.
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
"""
|
||||
Sample new images from a pre-trained DiT.
|
||||
"""
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
from diffusers.models import AutoencoderKL
|
||||
from torchvision.utils import save_image
|
||||
|
||||
from opendit.diffusion import create_diffusion
|
||||
from opendit.models.dit import DiT_models
|
||||
from opendit.utils.download import find_model
|
||||
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
torch.backends.cudnn.allow_tf32 = True
|
||||
|
||||
|
||||
def main(args):
|
||||
# Setup PyTorch:
|
||||
torch.manual_seed(args.seed)
|
||||
torch.set_grad_enabled(False)
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
if args.ckpt is None:
|
||||
raise ValueError("Please specify a checkpoint path with --ckpt.")
|
||||
|
||||
# Load model:
|
||||
vae = AutoencoderKL.from_pretrained(f"stabilityai/sd-vae-ft-{args.vae}").to(device)
|
||||
|
||||
# Configure input size
|
||||
assert args.image_size % 8 == 0, "Image size must be divisible by 8 (for the VAE encoder)."
|
||||
input_size = args.image_size // 8
|
||||
|
||||
dtype = torch.float32
|
||||
model = (
|
||||
DiT_models[args.model](
|
||||
input_size=input_size,
|
||||
num_classes=args.num_classes,
|
||||
enable_flashattn=False,
|
||||
enable_layernorm_kernel=False,
|
||||
dtype=dtype,
|
||||
)
|
||||
.to(device)
|
||||
.to(dtype)
|
||||
)
|
||||
|
||||
# Auto-download a pre-trained model or load a custom DiT checkpoint from train.py:
|
||||
ckpt_path = args.ckpt
|
||||
state_dict = find_model(ckpt_path)
|
||||
model.load_state_dict(state_dict)
|
||||
model.eval() # important!
|
||||
diffusion = create_diffusion(str(args.num_sampling_steps))
|
||||
|
||||
# Create sampling noise:
|
||||
# Labels to condition the model with (feel free to change):
|
||||
if args.num_classes == 1000:
|
||||
class_labels = [207, 360, 387, 974, 88, 979, 417, 279]
|
||||
else:
|
||||
class_labels = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9]
|
||||
n = len(class_labels)
|
||||
z = torch.randn(n, 4, input_size, input_size, device=device)
|
||||
y = torch.tensor(class_labels, device=device)
|
||||
y_null = torch.tensor([0] * n, device=device)
|
||||
y = torch.cat([y, y_null], 0)
|
||||
|
||||
# Setup classifier-free guidance:
|
||||
z = torch.cat([z, z], 0)
|
||||
model_kwargs = dict(y=y, cfg_scale=args.cfg_scale)
|
||||
|
||||
# Sample images:
|
||||
samples = diffusion.p_sample_loop(
|
||||
model.forward_with_cfg, z.shape, z, clip_denoised=False, model_kwargs=model_kwargs, progress=True, device=device
|
||||
)
|
||||
samples, _ = samples.chunk(2, dim=0) # Remove null class samples
|
||||
|
||||
# Save and display images:
|
||||
samples = vae.decode(samples / 0.18215).sample
|
||||
save_image(samples, "sample.png", nrow=4, normalize=True, value_range=(-1, 1))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--model", type=str, choices=DiT_models.keys(), default="DiT-XL/2")
|
||||
parser.add_argument("--vae", type=str, choices=["ema", "mse"], default="ema")
|
||||
parser.add_argument("--image_size", type=int, choices=[256, 512], default=256)
|
||||
parser.add_argument("--num_classes", type=int, default=1000)
|
||||
parser.add_argument("--cfg_scale", type=float, default=4.0)
|
||||
parser.add_argument("--num_sampling_steps", type=int, default=250)
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument(
|
||||
"--ckpt",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Optional path to a DiT checkpoint (default: auto-download a pre-trained DiT-XL/2 model).",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
Executable
+5
@@ -0,0 +1,5 @@
|
||||
python scripts/dit/sample_dit.py \
|
||||
--model DiT-XL/2 \
|
||||
--image_size 256 \
|
||||
--num_classes 1000 \
|
||||
--ckpt ckpt_path
|
||||
Executable
+324
@@ -0,0 +1,324 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
from glob import glob
|
||||
|
||||
import colossalai
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from colossalai.booster import Booster
|
||||
from colossalai.booster.plugin import LowLevelZeroPlugin, TorchDDPPlugin
|
||||
from colossalai.cluster import DistCoordinator
|
||||
from colossalai.nn.optimizer import HybridAdam
|
||||
from colossalai.utils import get_current_device
|
||||
from diffusers.models import AutoencoderKL
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
from torchvision.datasets import CIFAR10
|
||||
from tqdm import tqdm
|
||||
|
||||
from opendit.core.comm import model_sharding
|
||||
from opendit.core.parallel_mgr import get_parallel_manager, set_parallel_manager
|
||||
from opendit.datasets.dataloader import prepare_dataloader
|
||||
from opendit.datasets.image_transform import get_transforms_image
|
||||
from opendit.diffusion import create_diffusion
|
||||
from opendit.models.dit import DiT, DiT_models
|
||||
from opendit.utils.ckpt_utils import create_logger, load, record_model_param_shape, save
|
||||
from opendit.utils.train_utils import all_reduce_mean, format_numel_str, get_model_numel, requires_grad, update_ema
|
||||
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
torch.backends.cudnn.allow_tf32 = True
|
||||
|
||||
|
||||
def main(args):
|
||||
"""
|
||||
Trains a new DiT model.
|
||||
"""
|
||||
assert torch.cuda.is_available(), "Training currently requires at least one GPU."
|
||||
|
||||
# ==============================
|
||||
# Initialize Distributed Training
|
||||
# ==============================
|
||||
colossalai.launch_from_torch({}, seed=args.global_seed)
|
||||
coordinator = DistCoordinator()
|
||||
device = get_current_device()
|
||||
|
||||
# ==============================
|
||||
# Setup an experiment folder
|
||||
# ==============================
|
||||
# Make outputs folder (holds all experiment subfolders)
|
||||
os.makedirs(args.outputs, exist_ok=True)
|
||||
experiment_index = len(glob(f"{args.outputs}/*"))
|
||||
# e.g., DiT-XL/2 --> DiT-XL-2 (for naming folders)
|
||||
model_string_name = args.model.replace("/", "-")
|
||||
# Create an experiment folder
|
||||
experiment_dir = f"{args.outputs}/{experiment_index:03d}-{model_string_name}"
|
||||
dist.barrier()
|
||||
if coordinator.is_master():
|
||||
os.makedirs(experiment_dir, exist_ok=True)
|
||||
with open(f"{experiment_dir}/config.txt", "w") as f:
|
||||
json.dump(args.__dict__, f, indent=4)
|
||||
logger = create_logger(experiment_dir)
|
||||
logger.info(f"Experiment directory created at {experiment_dir}")
|
||||
else:
|
||||
logger = create_logger(None)
|
||||
|
||||
# ==============================
|
||||
# Initialize Tensorboard
|
||||
# ==============================
|
||||
if coordinator.is_master():
|
||||
tensorboard_dir = f"{experiment_dir}/tensorboard"
|
||||
os.makedirs(tensorboard_dir, exist_ok=True)
|
||||
writer = SummaryWriter(tensorboard_dir)
|
||||
|
||||
# ==============================
|
||||
# Initialize Booster
|
||||
# ==============================
|
||||
if args.plugin == "zero2":
|
||||
plugin = LowLevelZeroPlugin(
|
||||
stage=2,
|
||||
precision=args.mixed_precision,
|
||||
initial_scale=2**16,
|
||||
max_norm=args.grad_clip,
|
||||
)
|
||||
elif args.plugin == "ddp":
|
||||
plugin = TorchDDPPlugin()
|
||||
else:
|
||||
raise ValueError(f"Unknown plugin {args.plugin}")
|
||||
booster = Booster(plugin=plugin)
|
||||
|
||||
# ==============================
|
||||
# Initialize Process Group
|
||||
# ==============================
|
||||
sp_size = 1 # image doesn't need sequence parallel
|
||||
dp_size = dist.get_world_size() // sp_size
|
||||
set_parallel_manager(dp_size, sp_size, dp_axis=0, sp_axis=1)
|
||||
|
||||
# ======================================================
|
||||
# Initialize Model, Objective, Optimizer
|
||||
# ======================================================
|
||||
# Set mixed precision
|
||||
if args.mixed_precision == "bf16" and args.plugin != "ddp":
|
||||
dtype = torch.bfloat16
|
||||
elif args.mixed_precision == "fp16" and args.plugin != "ddp":
|
||||
dtype = torch.float16
|
||||
elif args.mixed_precision == "fp32" and args.plugin == "ddp":
|
||||
dtype = torch.float32
|
||||
else:
|
||||
raise ValueError(f"Unknown mixed precision {args.mixed_precision}")
|
||||
|
||||
# Create VAE encoder
|
||||
vae = AutoencoderKL.from_pretrained(f"stabilityai/sd-vae-ft-{args.vae}").to(device).to(dtype)
|
||||
|
||||
# Configure input size
|
||||
assert args.image_size % 8 == 0, "Image size must be divisible by 8 (for the VAE encoder)."
|
||||
input_size = args.image_size // 8
|
||||
|
||||
# Shared model config for two models
|
||||
model_config = {
|
||||
"input_size": input_size,
|
||||
"num_classes": args.num_classes,
|
||||
"enable_layernorm_kernel": args.enable_layernorm_kernel,
|
||||
"enable_modulate_kernel": args.enable_modulate_kernel,
|
||||
}
|
||||
|
||||
# Create DiT model
|
||||
model: DiT = (
|
||||
DiT_models[args.model](
|
||||
enable_flashattn=args.enable_flashattn,
|
||||
dtype=dtype,
|
||||
**model_config,
|
||||
)
|
||||
.to(device)
|
||||
.to(dtype)
|
||||
)
|
||||
|
||||
model_numel = get_model_numel(model)
|
||||
logger.info(f"Model params: {format_numel_str(model_numel)}")
|
||||
if args.grad_checkpoint:
|
||||
model.enable_gradient_checkpointing()
|
||||
|
||||
# Create ema and vae model
|
||||
# Note that parameter initialization is done within the DiT constructor
|
||||
# Create an EMA of the model for use after training
|
||||
ema = DiT_models[args.model](**model_config).to(device)
|
||||
ema = ema.to(torch.float32)
|
||||
ema.load_state_dict(model.state_dict())
|
||||
requires_grad(ema, False)
|
||||
ema_shape_dict = record_model_param_shape(ema)
|
||||
|
||||
# Create diffusion
|
||||
# default: 1000 steps, linear noise schedule
|
||||
diffusion = create_diffusion(timestep_respacing="")
|
||||
|
||||
# Setup optimizer
|
||||
# We used default Adam betas=(0.9, 0.999) and a constant learning rate of 1e-4 in our paper
|
||||
optimizer = HybridAdam(
|
||||
filter(lambda p: p.requires_grad, model.parameters()), lr=args.lr, weight_decay=0, adamw_mode=True
|
||||
)
|
||||
# You can use a lr scheduler if you want
|
||||
# Recommend if you continue training from a model
|
||||
lr_scheduler = None
|
||||
|
||||
# Prepare models for training
|
||||
# Ensure EMA is initialized with synced weights
|
||||
update_ema(ema, model, decay=0, sharded=False)
|
||||
# important! This enables embedding dropout for classifier-free guidance
|
||||
model.train()
|
||||
# EMA model should always be in eval mode
|
||||
ema.eval()
|
||||
|
||||
# Setup data:
|
||||
# master process goes first
|
||||
if not coordinator.is_master():
|
||||
dist.barrier()
|
||||
# To use ImageNet, you need to download it and:
|
||||
# from torchvision.datasets import ImageFolder
|
||||
# dataset = ImageFolder(args.data_path, transform=get_transforms_image(args.image_size))
|
||||
dataset = CIFAR10(args.data_path, transform=get_transforms_image(args.image_size), download=True)
|
||||
if coordinator.is_master():
|
||||
dist.barrier()
|
||||
dataloader = prepare_dataloader(
|
||||
dataset,
|
||||
batch_size=args.batch_size,
|
||||
shuffle=True,
|
||||
drop_last=True,
|
||||
pin_memory=True,
|
||||
num_workers=args.num_workers,
|
||||
pg_manager=get_parallel_manager(),
|
||||
)
|
||||
logger.info(f"Dataset contains {len(dataset):,} images ({args.data_path})")
|
||||
|
||||
# Boost model for distributed training
|
||||
torch.set_default_dtype(dtype)
|
||||
model, optimizer, _, dataloader, lr_scheduler = booster.boost(
|
||||
model=model, optimizer=optimizer, lr_scheduler=lr_scheduler, dataloader=dataloader
|
||||
)
|
||||
torch.set_default_dtype(torch.float)
|
||||
logger.info("Boost model for distributed training")
|
||||
|
||||
# Variables for monitoring/logging purposes:
|
||||
start_epoch = 0
|
||||
start_step = 0
|
||||
sampler_start_idx = 0
|
||||
if args.load is not None:
|
||||
logger.info("Loading checkpoint")
|
||||
start_epoch, start_step, sampler_start_idx = load(booster, model, ema, optimizer, lr_scheduler, args.load)
|
||||
logger.info(f"Loaded checkpoint {args.load} at epoch {start_epoch} step {start_step}")
|
||||
|
||||
# Only shard ema model when using zero2 plugin
|
||||
shard_ema = True if args.plugin == "zero2" else False
|
||||
if shard_ema:
|
||||
model_sharding(ema)
|
||||
|
||||
num_steps_per_epoch = len(dataloader)
|
||||
|
||||
logger.info(f"Training for {args.epochs} epochs...")
|
||||
# if resume training, set the sampler start index to the correct value
|
||||
dataloader.sampler.set_start_index(sampler_start_idx)
|
||||
for epoch in range(start_epoch, args.epochs):
|
||||
dataloader.sampler.set_epoch(epoch)
|
||||
dataloader_iter = iter(dataloader)
|
||||
logger.info(f"Beginning epoch {epoch}...")
|
||||
with tqdm(
|
||||
range(start_step, num_steps_per_epoch),
|
||||
desc=f"Epoch {epoch}",
|
||||
disable=not coordinator.is_master(),
|
||||
total=num_steps_per_epoch,
|
||||
initial=start_step,
|
||||
) as pbar:
|
||||
for step in pbar:
|
||||
x, y = next(dataloader_iter)
|
||||
x = x.to(device)
|
||||
y = y.to(device)
|
||||
|
||||
# VAE encode
|
||||
with torch.no_grad():
|
||||
# Map input images to latent space + normalize latents:
|
||||
x = x.to(dtype)
|
||||
x = vae.encode(x).latent_dist.sample().mul_(0.18215)
|
||||
# cast back to fp32 for bettet diffusion accuracy
|
||||
x = x.to(torch.float32)
|
||||
|
||||
# Diffusion
|
||||
t = torch.randint(0, diffusion.num_timesteps, (x.shape[0],), device=device)
|
||||
model_kwargs = dict(y=y)
|
||||
loss_dict = diffusion.training_losses(model, x, t, model_kwargs)
|
||||
loss = loss_dict["loss"].mean()
|
||||
booster.backward(loss=loss, optimizer=optimizer)
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
|
||||
# Update EMA
|
||||
update_ema(ema, model.unwrap(), optimizer=optimizer, sharded=shard_ema)
|
||||
|
||||
# Log loss values:
|
||||
all_reduce_mean(loss)
|
||||
global_step = epoch * num_steps_per_epoch + step
|
||||
pbar.set_postfix({"loss": loss.item(), "step": step, "global_step": global_step})
|
||||
|
||||
# Log to tensorboard
|
||||
if coordinator.is_master() and (global_step + 1) % args.log_every == 0:
|
||||
writer.add_scalar("loss", loss.item(), global_step)
|
||||
|
||||
# Save checkpoint
|
||||
if args.ckpt_every > 0 and (global_step + 1) % args.ckpt_every == 0:
|
||||
logger.info(f"Saving checkpoint...")
|
||||
save(
|
||||
booster,
|
||||
model,
|
||||
ema,
|
||||
optimizer,
|
||||
lr_scheduler,
|
||||
epoch,
|
||||
step + 1,
|
||||
global_step + 1,
|
||||
args.batch_size,
|
||||
coordinator,
|
||||
experiment_dir,
|
||||
ema_shape_dict,
|
||||
shard_ema,
|
||||
)
|
||||
logger.info(
|
||||
f"Saved checkpoint at epoch {epoch} step {step + 1} global_step {global_step + 1} to {experiment_dir}"
|
||||
)
|
||||
|
||||
# the continue epochs are not resumed, so we need to reset the sampler start index and start step
|
||||
dataloader.sampler.set_start_index(0)
|
||||
start_step = 0
|
||||
|
||||
model.eval() # important! This disables randomized embedding dropout
|
||||
# do any sampling/FID calculation/etc. with ema (or model) in eval mode ...
|
||||
|
||||
logger.info("Done!")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--model", type=str, choices=DiT_models.keys(), default="DiT-XL/2")
|
||||
parser.add_argument("--vae", type=str, choices=["ema", "mse"], default="ema") # Choice doesn't affect training
|
||||
parser.add_argument("--plugin", type=str, default="zero2")
|
||||
parser.add_argument("--outputs", type=str, default="./outputs", help="Path to the output directory")
|
||||
parser.add_argument("--load", type=str, default=None, help="Path to a checkpoint dir to load")
|
||||
|
||||
parser.add_argument("--data_path", type=str, default="./datasets", help="Path to the dataset")
|
||||
parser.add_argument("--image_size", type=int, choices=[256, 512], default=256)
|
||||
parser.add_argument("--num_classes", type=int, default=1000)
|
||||
|
||||
parser.add_argument("--epochs", type=int, default=1400)
|
||||
parser.add_argument("--batch_size", type=int, default=2)
|
||||
parser.add_argument("--global_seed", type=int, default=42)
|
||||
parser.add_argument("--num_workers", type=int, default=4)
|
||||
parser.add_argument("--log_every", type=int, default=10)
|
||||
parser.add_argument("--ckpt_every", type=int, default=1000)
|
||||
|
||||
parser.add_argument("--mixed_precision", type=str, default="bf16", choices=["bf16", "fp16", "fp32"])
|
||||
parser.add_argument("--grad_clip", type=float, default=1.0, help="Gradient clipping value")
|
||||
parser.add_argument("--lr", type=float, default=1e-4, help="Gradient clipping value")
|
||||
parser.add_argument("--grad_checkpoint", action="store_true", help="Use gradient checkpointing")
|
||||
|
||||
parser.add_argument("--enable_modulate_kernel", action="store_true", help="Enable triton modulate kernel")
|
||||
parser.add_argument("--enable_layernorm_kernel", action="store_true", help="Enable apex layernorm kernel")
|
||||
parser.add_argument("--enable_flashattn", action="store_true", help="Enable flashattn kernel")
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
Executable
+13
@@ -0,0 +1,13 @@
|
||||
torchrun --standalone --nproc_per_node=2 scripts/dit/train_dit.py \
|
||||
--model DiT-XL/2 \
|
||||
--batch_size 2 \
|
||||
--num_classes 10
|
||||
|
||||
# recommend setting
|
||||
# torchrun --standalone --nproc_per_node=8 scripts/dit/train_dit.py \
|
||||
# --model DiT-XL/2 \
|
||||
# --batch_size 180 \
|
||||
# --enable_layernorm_kernel \
|
||||
# --enable_flashattn \
|
||||
# --mixed_precision bf16 \
|
||||
# --num_classes 1000
|
||||
Executable
+262
@@ -0,0 +1,262 @@
|
||||
# Adapted from Latte
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Latte: https://github.com/Vchitect/Latte
|
||||
# --------------------------------------------------------
|
||||
|
||||
import argparse
|
||||
import os
|
||||
|
||||
import colossalai
|
||||
import imageio
|
||||
import torch
|
||||
from colossalai.cluster import DistCoordinator
|
||||
from diffusers.models import AutoencoderKL, AutoencoderKLTemporalDecoder
|
||||
from diffusers.schedulers import (
|
||||
DDIMScheduler,
|
||||
DDPMScheduler,
|
||||
DEISMultistepScheduler,
|
||||
DPMSolverMultistepScheduler,
|
||||
EulerAncestralDiscreteScheduler,
|
||||
EulerDiscreteScheduler,
|
||||
HeunDiscreteScheduler,
|
||||
KDPM2AncestralDiscreteScheduler,
|
||||
PNDMScheduler,
|
||||
)
|
||||
from diffusers.schedulers.scheduling_dpmsolver_singlestep import DPMSolverSinglestepScheduler
|
||||
from omegaconf import OmegaConf
|
||||
from torchvision.utils import save_image
|
||||
from transformers import T5EncoderModel, T5Tokenizer
|
||||
|
||||
from opendit.core.pab_mgr import set_pab_manager
|
||||
from opendit.core.parallel_mgr import set_parallel_manager
|
||||
from opendit.models.latte import LattePipeline, LatteT2V
|
||||
from opendit.utils.utils import merge_args, set_seed
|
||||
|
||||
|
||||
def main(args):
|
||||
set_seed(args.seed)
|
||||
torch.set_grad_enabled(False)
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
torch.backends.cudnn.allow_tf32 = True
|
||||
|
||||
# == init distributed env ==
|
||||
colossalai.launch_from_torch({})
|
||||
coordinator = DistCoordinator()
|
||||
set_parallel_manager(1, coordinator.world_size)
|
||||
device = f"cuda:{torch.cuda.current_device()}"
|
||||
|
||||
if args.cross_broadcast or args.spatial_broadcast or args.temporal_broadcast:
|
||||
set_pab_manager(
|
||||
steps=args.num_sampling_steps,
|
||||
cross_broadcast=args.cross_broadcast,
|
||||
cross_threshold=args.cross_threshold,
|
||||
cross_gap=args.cross_gap,
|
||||
spatial_broadcast=args.spatial_broadcast,
|
||||
spatial_threshold=args.spatial_threshold,
|
||||
spatial_gap=args.spatial_gap,
|
||||
temporal_broadcast=args.temporal_broadcast,
|
||||
temporal_threshold=args.temporal_threshold,
|
||||
temporal_gap=args.temporal_gap,
|
||||
diffusion_skip=args.diffusion_skip,
|
||||
diffusion_skip_timestep=args.diffusion_skip_timestep,
|
||||
)
|
||||
|
||||
transformer_model = LatteT2V.from_pretrained(
|
||||
args.pretrained_model_path, subfolder="transformer", video_length=args.video_length
|
||||
).to(device, dtype=torch.float16)
|
||||
|
||||
if args.enable_vae_temporal_decoder:
|
||||
vae = AutoencoderKLTemporalDecoder.from_pretrained(
|
||||
args.pretrained_model_path, subfolder="vae_temporal_decoder", torch_dtype=torch.float16
|
||||
).to(device)
|
||||
else:
|
||||
vae = AutoencoderKL.from_pretrained(args.pretrained_model_path, subfolder="vae", torch_dtype=torch.float16).to(
|
||||
device
|
||||
)
|
||||
tokenizer = T5Tokenizer.from_pretrained(args.pretrained_model_path, subfolder="tokenizer")
|
||||
text_encoder = T5EncoderModel.from_pretrained(
|
||||
args.pretrained_model_path, subfolder="text_encoder", torch_dtype=torch.float16
|
||||
).to(device)
|
||||
|
||||
# set eval mode
|
||||
transformer_model.eval()
|
||||
vae.eval()
|
||||
text_encoder.eval()
|
||||
|
||||
if args.sample_method == "DDIM":
|
||||
scheduler = DDIMScheduler.from_pretrained(
|
||||
args.pretrained_model_path,
|
||||
subfolder="scheduler",
|
||||
beta_start=args.beta_start,
|
||||
beta_end=args.beta_end,
|
||||
beta_schedule=args.beta_schedule,
|
||||
variance_type=args.variance_type,
|
||||
clip_sample=False,
|
||||
)
|
||||
elif args.sample_method == "EulerDiscrete":
|
||||
scheduler = EulerDiscreteScheduler.from_pretrained(
|
||||
args.pretrained_model_path,
|
||||
subfolder="scheduler",
|
||||
beta_start=args.beta_start,
|
||||
beta_end=args.beta_end,
|
||||
beta_schedule=args.beta_schedule,
|
||||
variance_type=args.variance_type,
|
||||
)
|
||||
elif args.sample_method == "DDPM":
|
||||
scheduler = DDPMScheduler.from_pretrained(
|
||||
args.pretrained_model_path,
|
||||
subfolder="scheduler",
|
||||
beta_start=args.beta_start,
|
||||
beta_end=args.beta_end,
|
||||
beta_schedule=args.beta_schedule,
|
||||
variance_type=args.variance_type,
|
||||
clip_sample=False,
|
||||
)
|
||||
elif args.sample_method == "DPMSolverMultistep":
|
||||
scheduler = DPMSolverMultistepScheduler.from_pretrained(
|
||||
args.pretrained_model_path,
|
||||
subfolder="scheduler",
|
||||
beta_start=args.beta_start,
|
||||
beta_end=args.beta_end,
|
||||
beta_schedule=args.beta_schedule,
|
||||
variance_type=args.variance_type,
|
||||
)
|
||||
elif args.sample_method == "DPMSolverSinglestep":
|
||||
scheduler = DPMSolverSinglestepScheduler.from_pretrained(
|
||||
args.pretrained_model_path,
|
||||
subfolder="scheduler",
|
||||
beta_start=args.beta_start,
|
||||
beta_end=args.beta_end,
|
||||
beta_schedule=args.beta_schedule,
|
||||
variance_type=args.variance_type,
|
||||
)
|
||||
elif args.sample_method == "PNDM":
|
||||
scheduler = PNDMScheduler.from_pretrained(
|
||||
args.pretrained_model_path,
|
||||
subfolder="scheduler",
|
||||
beta_start=args.beta_start,
|
||||
beta_end=args.beta_end,
|
||||
beta_schedule=args.beta_schedule,
|
||||
variance_type=args.variance_type,
|
||||
)
|
||||
elif args.sample_method == "HeunDiscrete":
|
||||
scheduler = HeunDiscreteScheduler.from_pretrained(
|
||||
args.pretrained_model_path,
|
||||
subfolder="scheduler",
|
||||
beta_start=args.beta_start,
|
||||
beta_end=args.beta_end,
|
||||
beta_schedule=args.beta_schedule,
|
||||
variance_type=args.variance_type,
|
||||
)
|
||||
elif args.sample_method == "EulerAncestralDiscrete":
|
||||
scheduler = EulerAncestralDiscreteScheduler.from_pretrained(
|
||||
args.pretrained_model_path,
|
||||
subfolder="scheduler",
|
||||
beta_start=args.beta_start,
|
||||
beta_end=args.beta_end,
|
||||
beta_schedule=args.beta_schedule,
|
||||
variance_type=args.variance_type,
|
||||
)
|
||||
elif args.sample_method == "DEISMultistep":
|
||||
scheduler = DEISMultistepScheduler.from_pretrained(
|
||||
args.pretrained_model_path,
|
||||
subfolder="scheduler",
|
||||
beta_start=args.beta_start,
|
||||
beta_end=args.beta_end,
|
||||
beta_schedule=args.beta_schedule,
|
||||
variance_type=args.variance_type,
|
||||
)
|
||||
elif args.sample_method == "KDPM2AncestralDiscrete":
|
||||
scheduler = KDPM2AncestralDiscreteScheduler.from_pretrained(
|
||||
args.pretrained_model_path,
|
||||
subfolder="scheduler",
|
||||
beta_start=args.beta_start,
|
||||
beta_end=args.beta_end,
|
||||
beta_schedule=args.beta_schedule,
|
||||
variance_type=args.variance_type,
|
||||
)
|
||||
|
||||
videogen_pipeline = LattePipeline(
|
||||
vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, scheduler=scheduler, transformer=transformer_model
|
||||
).to(device)
|
||||
|
||||
os.makedirs(args.save_img_path, exist_ok=True)
|
||||
|
||||
# video_grids = []
|
||||
for num_prompt, prompt in enumerate(args.text_prompt):
|
||||
print("Processing the ({}) prompt".format(prompt))
|
||||
videos = videogen_pipeline(
|
||||
prompt,
|
||||
video_length=args.video_length,
|
||||
height=args.image_size[0],
|
||||
width=args.image_size[1],
|
||||
num_inference_steps=args.num_sampling_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
enable_temporal_attentions=args.enable_temporal_attentions,
|
||||
num_images_per_prompt=1,
|
||||
mask_feature=True,
|
||||
enable_vae_temporal_decoder=args.enable_vae_temporal_decoder,
|
||||
).video
|
||||
if coordinator.is_master():
|
||||
if videos.shape[1] == 1:
|
||||
save_image(videos[0][0], args.save_img_path + prompt[:30].replace(" ", "_") + ".png")
|
||||
else:
|
||||
imageio.mimwrite(
|
||||
args.save_img_path + prompt[:30].replace(" ", "_") + "_%04d" % args.run_time + ".mp4",
|
||||
videos[0],
|
||||
fps=8,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--config", type=str, required=True)
|
||||
parser.add_argument("--save_img_path", type=str, default="./samples/latte/")
|
||||
parser.add_argument("--pretrained_model_path", type=str, default="maxin-cn/Latte-1")
|
||||
parser.add_argument("--model", type=str, default="LatteT2V")
|
||||
parser.add_argument("--video_length", type=int, default=16)
|
||||
parser.add_argument("--image_size", nargs="+")
|
||||
parser.add_argument("--beta_start", type=float, default=0.0001)
|
||||
parser.add_argument("--beta_end", type=float, default=0.02)
|
||||
parser.add_argument("--beta_schedule", type=str, default="linear")
|
||||
parser.add_argument("--variance_type", type=str, default="learned_range")
|
||||
parser.add_argument("--use_compile", action="store_true")
|
||||
parser.add_argument("--use_fp16", action="store_true")
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--run_time", type=int, default=0)
|
||||
parser.add_argument("--guidance_scale", type=float, default=7.5)
|
||||
parser.add_argument("--sample_method", type=str, default="DDIM")
|
||||
parser.add_argument("--num_sampling_steps", type=int, default=50)
|
||||
parser.add_argument("--enable_temporal_attentions", action="store_true")
|
||||
parser.add_argument("--enable_vae_temporal_decoder", action="store_true")
|
||||
parser.add_argument("--text_prompt", nargs="+")
|
||||
|
||||
# pab
|
||||
parser.add_argument("--spatial_broadcast", action="store_true", help="Enable spatial attention skip")
|
||||
parser.add_argument(
|
||||
"--spatial_threshold", type=int, nargs=2, default=[100, 800], help="Spatial attention threshold"
|
||||
)
|
||||
parser.add_argument("--spatial_gap", type=int, default=2, help="Spatial attention gap")
|
||||
parser.add_argument("--temporal_broadcast", action="store_true", help="Enable temporal attention skip")
|
||||
parser.add_argument(
|
||||
"--temporal_threshold", type=int, nargs=2, default=[100, 800], help="Temporal attention threshold"
|
||||
)
|
||||
parser.add_argument("--temporal_gap", type=int, default=4, help="Temporal attention gap")
|
||||
parser.add_argument("--cross_broadcast", action="store_true", help="Enable cross attention skip")
|
||||
parser.add_argument("--cross_threshold", type=int, nargs=2, default=[80, 900], help="Cross attention threshold")
|
||||
parser.add_argument("--cross_gap", type=int, default=7, help="Cross attention gap")
|
||||
parser.add_argument(
|
||||
"--diffusion_skip",
|
||||
action="store_true",
|
||||
)
|
||||
parser.add_argument("--diffusion_skip_timestep", nargs="+")
|
||||
|
||||
args = parser.parse_args()
|
||||
config_args = OmegaConf.load(args.config)
|
||||
args = merge_args(args, config_args)
|
||||
|
||||
main(args)
|
||||
Executable
+1
@@ -0,0 +1 @@
|
||||
torchrun --standalone --nproc_per_node=1 scripts/latte/sample.py --config configs/latte/sample.yaml
|
||||
Executable
+2
@@ -0,0 +1,2 @@
|
||||
# Pyramidal Attention Broadcast
|
||||
torchrun --standalone --nproc_per_node=8 scripts/latte/sample.py --config configs/latte/sample_pab.yaml
|
||||
Executable
+384
@@ -0,0 +1,384 @@
|
||||
import argparse
|
||||
import os
|
||||
import time
|
||||
|
||||
import colossalai
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from colossalai.cluster import DistCoordinator
|
||||
from omegaconf import OmegaConf
|
||||
from tqdm import tqdm
|
||||
|
||||
from opendit.core.pab_mgr import set_pab_manager
|
||||
from opendit.core.parallel_mgr import enable_sequence_parallel, set_parallel_manager
|
||||
from opendit.models.opensora import RFLOW, OpenSoraVAE_V1_2, STDiT3_XL_2, T5Encoder, text_preprocessing
|
||||
from opendit.models.opensora.datasets import get_image_size, get_num_frames, save_sample
|
||||
from opendit.models.opensora.inference_utils import (
|
||||
add_watermark,
|
||||
append_generated,
|
||||
append_score_to_prompts,
|
||||
apply_mask_strategy,
|
||||
collect_references_batch,
|
||||
dframe_to_frame,
|
||||
extract_json_from_prompts,
|
||||
extract_prompts_loop,
|
||||
get_save_path_name,
|
||||
load_prompts,
|
||||
merge_prompt,
|
||||
prepare_multi_resolution_info,
|
||||
refine_prompts_by_openai,
|
||||
split_prompt,
|
||||
)
|
||||
from opendit.utils.utils import all_exists, create_logger, merge_args, set_seed, str_to_dtype
|
||||
|
||||
|
||||
def main(args):
|
||||
torch.set_grad_enabled(False)
|
||||
# ======================================================
|
||||
# configs & runtime variables
|
||||
# ======================================================
|
||||
# == dtype ==
|
||||
dtype = str_to_dtype(args.dtype)
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
torch.backends.cudnn.allow_tf32 = True
|
||||
|
||||
# == init distributed env ==
|
||||
colossalai.launch_from_torch({})
|
||||
coordinator = DistCoordinator()
|
||||
set_parallel_manager(1, coordinator.world_size)
|
||||
enable_sequence_parallelism = enable_sequence_parallel()
|
||||
device = f"cuda:{torch.cuda.current_device()}"
|
||||
set_seed(seed=args.seed)
|
||||
|
||||
# == init pab ==
|
||||
if args.cross_broadcast or args.spatial_broadcast or args.temporal_broadcast:
|
||||
set_pab_manager(
|
||||
steps=args.num_sampling_steps,
|
||||
cross_broadcast=args.cross_broadcast,
|
||||
cross_threshold=args.cross_threshold,
|
||||
cross_gap=args.cross_gap,
|
||||
spatial_broadcast=args.spatial_broadcast,
|
||||
spatial_threshold=args.spatial_threshold,
|
||||
spatial_gap=args.spatial_gap,
|
||||
temporal_broadcast=args.temporal_broadcast,
|
||||
temporal_threshold=args.temporal_threshold,
|
||||
temporal_gap=args.temporal_gap,
|
||||
diffusion_skip=args.diffusion_skip,
|
||||
diffusion_skip_timestep=args.diffusion_skip_timestep,
|
||||
)
|
||||
|
||||
# == init logger ==
|
||||
logger = create_logger()
|
||||
logger.info(f"Inference configuration: {args}\n")
|
||||
verbose = args.verbose
|
||||
progress_wrap = tqdm if verbose == 1 else (lambda x: x)
|
||||
|
||||
# ======================================================
|
||||
# build model & load weights
|
||||
# ======================================================
|
||||
logger.info("Building models...")
|
||||
# == build text-encoder and vae ==
|
||||
text_encoder = T5Encoder(
|
||||
from_pretrained="DeepFloyd/t5-v1_1-xxl", model_max_length=300, device=device, shardformer=args.enable_t5_speedup
|
||||
)
|
||||
vae = (
|
||||
OpenSoraVAE_V1_2(
|
||||
from_pretrained="hpcai-tech/OpenSora-VAE-v1.2",
|
||||
micro_frame_size=17,
|
||||
micro_batch_size=4,
|
||||
)
|
||||
.to(device, dtype)
|
||||
.eval()
|
||||
)
|
||||
|
||||
# == prepare video size ==
|
||||
image_size = args.image_size
|
||||
if image_size is None:
|
||||
resolution = args.resolution
|
||||
aspect_ratio = args.aspect_ratio
|
||||
assert (
|
||||
resolution is not None and aspect_ratio is not None
|
||||
), "resolution and aspect_ratio must be provided if image_size is not provided"
|
||||
image_size = get_image_size(resolution, aspect_ratio)
|
||||
num_frames = get_num_frames(args.num_frames)
|
||||
|
||||
# == build diffusion model ==
|
||||
input_size = (num_frames, *image_size)
|
||||
latent_size = vae.get_latent_size(input_size)
|
||||
model = (
|
||||
STDiT3_XL_2(
|
||||
from_pretrained="hpcai-tech/OpenSora-STDiT-v3",
|
||||
qk_norm=True,
|
||||
enable_flash_attn=True,
|
||||
enable_layernorm_kernel=True,
|
||||
input_size=latent_size,
|
||||
in_channels=vae.out_channels,
|
||||
caption_channels=text_encoder.output_dim,
|
||||
model_max_length=text_encoder.model_max_length,
|
||||
)
|
||||
.to(device, dtype)
|
||||
.eval()
|
||||
)
|
||||
text_encoder.y_embedder = model.y_embedder # HACK: for classifier-free guidance
|
||||
|
||||
# == build scheduler ==
|
||||
scheduler = RFLOW(use_timestep_transform=True, num_sampling_steps=30, cfg_scale=7.0)
|
||||
|
||||
# ======================================================
|
||||
# inference
|
||||
# ======================================================
|
||||
# == load prompts ==
|
||||
prompts = args.prompt
|
||||
if prompts is None:
|
||||
assert args.prompt_path is not None
|
||||
prompts = load_prompts(args.prompt_path)
|
||||
|
||||
# == prepare reference ==
|
||||
reference_path = args.reference_path if args.reference_path is not None else [""] * len(prompts)
|
||||
mask_strategy = args.mask_strategy if args.mask_strategy is not None else [""] * len(prompts)
|
||||
assert len(reference_path) == len(prompts), "Length of reference must be the same as prompts"
|
||||
assert len(mask_strategy) == len(prompts), "Length of mask_strategy must be the same as prompts"
|
||||
|
||||
# == prepare arguments ==
|
||||
fps = args.fps
|
||||
save_fps = fps // args.frame_interval
|
||||
multi_resolution = args.multi_resolution
|
||||
batch_size = args.batch_size
|
||||
num_sample = args.num_sample
|
||||
loop = args.loop
|
||||
condition_frame_length = args.condition_frame_length
|
||||
condition_frame_edit = args.condition_frame_edit
|
||||
align = args.align
|
||||
|
||||
save_dir = args.save_dir
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
prompt_as_path = args.prompt_as_path
|
||||
|
||||
# == Iter over all samples ==
|
||||
for i in progress_wrap(range(0, len(prompts), batch_size)):
|
||||
# == prepare batch prompts ==
|
||||
batch_prompts = prompts[i : i + batch_size]
|
||||
ms = mask_strategy[i : i + batch_size]
|
||||
refs = reference_path[i : i + batch_size]
|
||||
|
||||
# == get json from prompts ==
|
||||
batch_prompts, refs, ms = extract_json_from_prompts(batch_prompts, refs, ms)
|
||||
original_batch_prompts = batch_prompts
|
||||
|
||||
# == get reference for condition ==
|
||||
refs = collect_references_batch(refs, vae, image_size)
|
||||
|
||||
# == multi-resolution info ==
|
||||
model_args = prepare_multi_resolution_info(
|
||||
multi_resolution, len(batch_prompts), image_size, num_frames, fps, device, dtype
|
||||
)
|
||||
|
||||
# == Iter over number of sampling for one prompt ==
|
||||
for k in range(num_sample):
|
||||
# == prepare save paths ==
|
||||
save_paths = [
|
||||
get_save_path_name(
|
||||
save_dir,
|
||||
sample_idx=idx,
|
||||
prompt=original_batch_prompts[idx],
|
||||
prompt_as_path=prompt_as_path,
|
||||
num_sample=num_sample,
|
||||
k=k,
|
||||
)
|
||||
for idx in range(len(batch_prompts))
|
||||
]
|
||||
|
||||
# NOTE: Skip if the sample already exists
|
||||
# This is useful for resuming sampling VBench
|
||||
if prompt_as_path and all_exists(save_paths):
|
||||
continue
|
||||
|
||||
# == process prompts step by step ==
|
||||
# 0. split prompt
|
||||
# each element in the list is [prompt_segment_list, loop_idx_list]
|
||||
batched_prompt_segment_list = []
|
||||
batched_loop_idx_list = []
|
||||
for prompt in batch_prompts:
|
||||
prompt_segment_list, loop_idx_list = split_prompt(prompt)
|
||||
batched_prompt_segment_list.append(prompt_segment_list)
|
||||
batched_loop_idx_list.append(loop_idx_list)
|
||||
|
||||
# 1. refine prompt by openai
|
||||
if args.llm_refine:
|
||||
# only call openai API when
|
||||
# 1. seq parallel is not enabled
|
||||
# 2. seq parallel is enabled and the process is rank 0
|
||||
if not enable_sequence_parallelism or (enable_sequence_parallelism and coordinator.is_master()):
|
||||
for idx, prompt_segment_list in enumerate(batched_prompt_segment_list):
|
||||
batched_prompt_segment_list[idx] = refine_prompts_by_openai(prompt_segment_list)
|
||||
|
||||
# sync the prompt if using seq parallel
|
||||
if enable_sequence_parallelism:
|
||||
coordinator.block_all()
|
||||
prompt_segment_length = [
|
||||
len(prompt_segment_list) for prompt_segment_list in batched_prompt_segment_list
|
||||
]
|
||||
|
||||
# flatten the prompt segment list
|
||||
batched_prompt_segment_list = [
|
||||
prompt_segment
|
||||
for prompt_segment_list in batched_prompt_segment_list
|
||||
for prompt_segment in prompt_segment_list
|
||||
]
|
||||
|
||||
# create a list of size equal to world size
|
||||
broadcast_obj_list = [batched_prompt_segment_list] * coordinator.world_size
|
||||
dist.broadcast_object_list(broadcast_obj_list, 0)
|
||||
|
||||
# recover the prompt list
|
||||
batched_prompt_segment_list = []
|
||||
segment_start_idx = 0
|
||||
all_prompts = broadcast_obj_list[0]
|
||||
for num_segment in prompt_segment_length:
|
||||
batched_prompt_segment_list.append(
|
||||
all_prompts[segment_start_idx : segment_start_idx + num_segment]
|
||||
)
|
||||
segment_start_idx += num_segment
|
||||
|
||||
# 2. append score
|
||||
for idx, prompt_segment_list in enumerate(batched_prompt_segment_list):
|
||||
batched_prompt_segment_list[idx] = append_score_to_prompts(
|
||||
prompt_segment_list,
|
||||
aes=args.aes,
|
||||
flow=args.flow,
|
||||
camera_motion=args.camera_motion,
|
||||
)
|
||||
|
||||
# 3. clean prompt with T5
|
||||
for idx, prompt_segment_list in enumerate(batched_prompt_segment_list):
|
||||
batched_prompt_segment_list[idx] = [text_preprocessing(prompt) for prompt in prompt_segment_list]
|
||||
|
||||
# 4. merge to obtain the final prompt
|
||||
batch_prompts = []
|
||||
for prompt_segment_list, loop_idx_list in zip(batched_prompt_segment_list, batched_loop_idx_list):
|
||||
batch_prompts.append(merge_prompt(prompt_segment_list, loop_idx_list))
|
||||
|
||||
# == Iter over loop generation ==
|
||||
video_clips = []
|
||||
for loop_i in range(loop):
|
||||
# == get prompt for loop i ==
|
||||
batch_prompts_loop = extract_prompts_loop(batch_prompts, loop_i)
|
||||
|
||||
# == add condition frames for loop ==
|
||||
if loop_i > 0:
|
||||
refs, ms = append_generated(
|
||||
vae, video_clips[-1], refs, ms, loop_i, condition_frame_length, condition_frame_edit
|
||||
)
|
||||
|
||||
# == sampling ==
|
||||
z = torch.randn(len(batch_prompts), vae.out_channels, *latent_size, device=device, dtype=dtype)
|
||||
masks = apply_mask_strategy(z, refs, ms, loop_i, align=align)
|
||||
samples = scheduler.sample(
|
||||
model,
|
||||
text_encoder,
|
||||
z=z,
|
||||
prompts=batch_prompts_loop,
|
||||
device=device,
|
||||
additional_args=model_args,
|
||||
progress=verbose >= 2,
|
||||
mask=masks,
|
||||
)
|
||||
samples = vae.decode(samples.to(dtype), num_frames=num_frames)
|
||||
video_clips.append(samples)
|
||||
|
||||
# == save samples ==
|
||||
if coordinator.is_master():
|
||||
for idx, batch_prompt in enumerate(batch_prompts):
|
||||
if verbose >= 2:
|
||||
logger.info("Prompt: %s", batch_prompt)
|
||||
save_path = save_paths[idx]
|
||||
video = [video_clips[i][idx] for i in range(loop)]
|
||||
for i in range(1, loop):
|
||||
video[i] = video[i][:, dframe_to_frame(condition_frame_length) :]
|
||||
video = torch.cat(video, dim=1)
|
||||
save_path = save_sample(
|
||||
video,
|
||||
fps=save_fps,
|
||||
save_path=save_path,
|
||||
verbose=verbose >= 2,
|
||||
)
|
||||
if save_path.endswith(".mp4") and args.watermark:
|
||||
time.sleep(1) # prevent loading previous generated video
|
||||
add_watermark(save_path)
|
||||
logger.info("Inference finished.")
|
||||
logger.info("Saved samples to %s", save_dir)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
# general
|
||||
parser.add_argument("--config", default=None, type=str, help="path to config yaml")
|
||||
parser.add_argument("--seed", default=1024, type=int, help="seed for reproducibility")
|
||||
parser.add_argument("--batch-size", default=1, type=int, help="batch size")
|
||||
parser.add_argument("--flash-attn", action="store_true", help="enable flash attention")
|
||||
parser.add_argument("--enable_t5_speedup", action="store_true", help="enable t5 speedup")
|
||||
parser.add_argument("--resolution", default=None, type=str, help="resolution")
|
||||
parser.add_argument("--multi-resolution", default=None, type=str, help="multi resolution")
|
||||
parser.add_argument("--dtype", default="bf16", type=str, help="data type")
|
||||
|
||||
# output
|
||||
parser.add_argument("--save-dir", default="./samples/opensora", type=str, help="path to save generated samples")
|
||||
parser.add_argument("--num-sample", default=1, type=int, help="number of samples to generate for one prompt")
|
||||
parser.add_argument("--prompt-as-path", action="store_true", help="use prompt as path to save samples")
|
||||
parser.add_argument("--verbose", default=2, type=int, help="verbose level")
|
||||
|
||||
# prompt
|
||||
parser.add_argument("--prompt-path", default=None, type=str, help="path to prompt txt file")
|
||||
parser.add_argument("--prompt", default=None, type=str, nargs="+", help="prompt list")
|
||||
parser.add_argument("--llm-refine", action="store_true", help="enable LLM refine")
|
||||
|
||||
# image/video
|
||||
parser.add_argument("--num-frames", default=None, type=str, help="number of frames")
|
||||
parser.add_argument("--fps", default=24, type=int, help="fps")
|
||||
parser.add_argument("--image-size", default=None, type=int, nargs=2, help="image size")
|
||||
parser.add_argument("--frame-interval", default=1, type=int, help="frame interval")
|
||||
parser.add_argument("--aspect-ratio", default=None, type=str, help="aspect ratio (h:w)")
|
||||
parser.add_argument("--watermark", action="store_true", help="watermark video")
|
||||
|
||||
# hyperparameters
|
||||
parser.add_argument("--num-sampling-steps", default=30, type=int, help="sampling steps")
|
||||
parser.add_argument("--cfg-scale", default=7.0, type=float, help="balance between cond & uncond")
|
||||
|
||||
# reference
|
||||
parser.add_argument("--loop", default=1, type=int, help="loop")
|
||||
parser.add_argument("--align", default=None, type=int, help="align")
|
||||
parser.add_argument("--condition-frame-length", default=5, type=int, help="condition frame length")
|
||||
parser.add_argument("--condition-frame-edit", default=0.0, type=float, help="condition frame edit")
|
||||
parser.add_argument("--reference-path", default=None, type=str, nargs="+", help="reference path")
|
||||
parser.add_argument("--mask-strategy", default=None, type=str, nargs="+", help="mask strategy")
|
||||
parser.add_argument("--aes", default=None, type=float, help="aesthetic score")
|
||||
parser.add_argument("--flow", default=None, type=float, help="flow score")
|
||||
parser.add_argument("--camera-motion", default=None, type=str, help="camera motion")
|
||||
|
||||
# pab
|
||||
parser.add_argument("--spatial_broadcast", action="store_true", help="Enable spatial attention skip")
|
||||
parser.add_argument(
|
||||
"--spatial_threshold", type=int, nargs=2, default=[540, 920], help="Spatial attention threshold"
|
||||
)
|
||||
parser.add_argument("--spatial_gap", type=int, default=2, help="Spatial attention gap")
|
||||
parser.add_argument("--temporal_broadcast", action="store_true", help="Enable temporal attention skip")
|
||||
parser.add_argument(
|
||||
"--temporal_threshold", type=int, nargs=2, default=[540, 960], help="Temporal attention threshold"
|
||||
)
|
||||
parser.add_argument("--temporal_gap", type=int, default=4, help="Temporal attention gap")
|
||||
parser.add_argument("--cross_broadcast", action="store_true", help="Enable cross attention skip")
|
||||
parser.add_argument("--cross_threshold", type=int, nargs=2, default=[540, 960], help="Cross attention threshold")
|
||||
parser.add_argument("--cross_gap", type=int, default=6, help="Cross attention gap")
|
||||
parser.add_argument(
|
||||
"--diffusion_skip",
|
||||
action="store_true",
|
||||
)
|
||||
parser.add_argument("--diffusion_skip_timestep", nargs="+")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
config_args = OmegaConf.load(args.config)
|
||||
args = merge_args(args, config_args)
|
||||
|
||||
main(args)
|
||||
Executable
+1
@@ -0,0 +1 @@
|
||||
torchrun --standalone --nproc_per_node=1 scripts/opensora/sample.py --config configs/opensora/sample.yaml
|
||||
Executable
+2
@@ -0,0 +1,2 @@
|
||||
# Pyramid Attention Broadcast
|
||||
torchrun --standalone --nproc_per_node=8 scripts/opensora/sample.py --config configs/opensora/sample_pab.yaml
|
||||
Executable
+258
@@ -0,0 +1,258 @@
|
||||
# Adapted from Open-Sora-Plan
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
# References:
|
||||
# Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
|
||||
# --------------------------------------------------------
|
||||
|
||||
import argparse
|
||||
import math
|
||||
import os
|
||||
|
||||
import colossalai
|
||||
import imageio
|
||||
import torch
|
||||
from colossalai.cluster import DistCoordinator
|
||||
from diffusers.schedulers import (
|
||||
DDIMScheduler,
|
||||
DDPMScheduler,
|
||||
DEISMultistepScheduler,
|
||||
DPMSolverMultistepScheduler,
|
||||
EulerAncestralDiscreteScheduler,
|
||||
EulerDiscreteScheduler,
|
||||
HeunDiscreteScheduler,
|
||||
KDPM2AncestralDiscreteScheduler,
|
||||
PNDMScheduler,
|
||||
)
|
||||
from diffusers.schedulers.scheduling_dpmsolver_singlestep import DPMSolverSinglestepScheduler
|
||||
from omegaconf import OmegaConf
|
||||
from torchvision.utils import save_image
|
||||
from transformers import T5EncoderModel, T5Tokenizer
|
||||
|
||||
from opendit.core.pab_mgr import set_pab_manager
|
||||
from opendit.core.parallel_mgr import set_parallel_manager
|
||||
from opendit.models.opensora_plan import LatteT2V, VideoGenPipeline, ae_stride_config, getae_wrapper
|
||||
from opendit.utils.utils import merge_args, set_seed
|
||||
|
||||
|
||||
def save_video_grid(video, nrow=None):
|
||||
b, t, h, w, c = video.shape
|
||||
|
||||
if nrow is None:
|
||||
nrow = math.ceil(math.sqrt(b))
|
||||
ncol = math.ceil(b / nrow)
|
||||
padding = 1
|
||||
video_grid = torch.zeros((t, (padding + h) * nrow + padding, (padding + w) * ncol + padding, c), dtype=torch.uint8)
|
||||
|
||||
for i in range(b):
|
||||
r = i // ncol
|
||||
c = i % ncol
|
||||
start_r = (padding + h) * r
|
||||
start_c = (padding + w) * c
|
||||
video_grid[:, start_r : start_r + h, start_c : start_c + w] = video[i]
|
||||
|
||||
return video_grid
|
||||
|
||||
|
||||
def main(args):
|
||||
set_seed(42)
|
||||
torch.set_grad_enabled(False)
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
torch.backends.cudnn.allow_tf32 = True
|
||||
|
||||
# == init distributed env ==
|
||||
colossalai.launch_from_torch({})
|
||||
coordinator = DistCoordinator()
|
||||
set_parallel_manager(1, coordinator.world_size)
|
||||
device = f"cuda:{torch.cuda.current_device()}"
|
||||
|
||||
if args.cross_broadcast or args.spatial_broadcast or args.temporal_broadcast:
|
||||
set_pab_manager(
|
||||
steps=args.num_sampling_steps,
|
||||
cross_broadcast=args.cross_broadcast,
|
||||
cross_threshold=args.cross_threshold,
|
||||
cross_gap=args.cross_gap,
|
||||
spatial_broadcast=args.spatial_broadcast,
|
||||
spatial_threshold=args.spatial_threshold,
|
||||
spatial_gap=args.spatial_gap,
|
||||
temporal_broadcast=args.temporal_broadcast,
|
||||
temporal_threshold=args.temporal_threshold,
|
||||
temporal_gap=args.temporal_gap,
|
||||
diffusion_skip=args.diffusion_skip,
|
||||
diffusion_skip_timestep=args.diffusion_skip_timestep,
|
||||
)
|
||||
|
||||
vae = getae_wrapper(args.ae)(args.model_path, subfolder="vae", cache_dir=args.cache_dir).to(
|
||||
device, dtype=torch.float16
|
||||
)
|
||||
# vae = getae_wrapper(args.ae)(args.ae_path).to(device, dtype=torch.float16)
|
||||
if args.enable_tiling:
|
||||
vae.vae.enable_tiling()
|
||||
vae.vae.tile_overlap_factor = args.tile_overlap_factor
|
||||
vae.vae_scale_factor = ae_stride_config[args.ae]
|
||||
# Load model:
|
||||
transformer_model = LatteT2V.from_pretrained(
|
||||
args.model_path, subfolder=args.version, cache_dir=args.cache_dir, torch_dtype=torch.float16
|
||||
).to(device)
|
||||
# transformer_model = LatteT2V.from_pretrained(args.model_path, low_cpu_mem_usage=False, device_map=None, torch_dtype=torch.float16).to(device)
|
||||
|
||||
transformer_model.force_images = args.force_images
|
||||
tokenizer = T5Tokenizer.from_pretrained(args.text_encoder_name, cache_dir=args.cache_dir)
|
||||
text_encoder = T5EncoderModel.from_pretrained(
|
||||
args.text_encoder_name, cache_dir=args.cache_dir, torch_dtype=torch.float16
|
||||
).to(device)
|
||||
|
||||
if args.force_images:
|
||||
ext = "jpg"
|
||||
else:
|
||||
ext = "mp4"
|
||||
|
||||
# set eval mode
|
||||
transformer_model.eval()
|
||||
vae.eval()
|
||||
text_encoder.eval()
|
||||
|
||||
if args.sample_method == "DDIM": #########
|
||||
scheduler = DDIMScheduler()
|
||||
elif args.sample_method == "EulerDiscrete":
|
||||
scheduler = EulerDiscreteScheduler()
|
||||
elif args.sample_method == "DDPM": #############
|
||||
scheduler = DDPMScheduler()
|
||||
elif args.sample_method == "DPMSolverMultistep":
|
||||
scheduler = DPMSolverMultistepScheduler()
|
||||
elif args.sample_method == "DPMSolverSinglestep":
|
||||
scheduler = DPMSolverSinglestepScheduler()
|
||||
elif args.sample_method == "PNDM":
|
||||
scheduler = PNDMScheduler()
|
||||
elif args.sample_method == "HeunDiscrete": ########
|
||||
scheduler = HeunDiscreteScheduler()
|
||||
elif args.sample_method == "EulerAncestralDiscrete":
|
||||
scheduler = EulerAncestralDiscreteScheduler()
|
||||
elif args.sample_method == "DEISMultistep":
|
||||
scheduler = DEISMultistepScheduler()
|
||||
elif args.sample_method == "KDPM2AncestralDiscrete": #########
|
||||
scheduler = KDPM2AncestralDiscreteScheduler()
|
||||
videogen_pipeline = VideoGenPipeline(
|
||||
vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, scheduler=scheduler, transformer=transformer_model
|
||||
).to(device=device)
|
||||
# videogen_pipeline.enable_xformers_memory_efficient_attention()
|
||||
|
||||
os.makedirs(args.save_img_path, exist_ok=True)
|
||||
|
||||
video_grids = []
|
||||
if not isinstance(args.text_prompt, list):
|
||||
args.text_prompt = [args.text_prompt]
|
||||
if len(args.text_prompt) == 1 and args.text_prompt[0].endswith("txt"):
|
||||
text_prompt = open(args.text_prompt[0], "r").readlines()
|
||||
args.text_prompt = [i.strip() for i in text_prompt]
|
||||
for idx, prompt in enumerate(args.text_prompt):
|
||||
print("Processing the ({}) prompt".format(prompt))
|
||||
videos = videogen_pipeline(
|
||||
prompt,
|
||||
num_frames=args.num_frames,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_inference_steps=args.num_sampling_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
enable_temporal_attentions=not args.force_images,
|
||||
num_images_per_prompt=1,
|
||||
mask_feature=True,
|
||||
).video
|
||||
try:
|
||||
if args.force_images:
|
||||
videos = videos[:, 0].permute(0, 3, 1, 2) # b t h w c -> b c h w
|
||||
save_image(
|
||||
videos / 255.0,
|
||||
os.path.join(args.save_img_path, f"{idx}.{ext}"),
|
||||
nrow=1,
|
||||
normalize=True,
|
||||
value_range=(0, 1),
|
||||
) # t c h w
|
||||
|
||||
else:
|
||||
imageio.mimwrite(
|
||||
os.path.join(args.save_img_path, f"{idx}.{ext}"), videos[0], fps=args.fps, quality=9
|
||||
) # highest quality is 10, lowest is 0
|
||||
except:
|
||||
print("Error when saving {}".format(prompt))
|
||||
video_grids.append(videos)
|
||||
video_grids = torch.cat(video_grids, dim=0)
|
||||
|
||||
# torchvision.io.write_video(args.save_img_path + '_%04d' % args.run_time + '-.mp4', video_grids, fps=6)
|
||||
if coordinator.is_master():
|
||||
if args.force_images:
|
||||
save_image(
|
||||
video_grids / 255.0,
|
||||
os.path.join(
|
||||
args.save_img_path, f"{args.sample_method}_gs{args.guidance_scale}_s{args.num_sampling_steps}.{ext}"
|
||||
),
|
||||
nrow=math.ceil(math.sqrt(len(video_grids))),
|
||||
normalize=True,
|
||||
value_range=(0, 1),
|
||||
)
|
||||
else:
|
||||
video_grids = save_video_grid(video_grids)
|
||||
imageio.mimwrite(
|
||||
os.path.join(
|
||||
args.save_img_path, f"{args.sample_method}_gs{args.guidance_scale}_s{args.num_sampling_steps}.{ext}"
|
||||
),
|
||||
video_grids,
|
||||
fps=args.fps,
|
||||
quality=9,
|
||||
)
|
||||
|
||||
print("save path {}".format(args.save_img_path))
|
||||
|
||||
# save_videos_grid(video, f"./{prompt}.gif")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--config", type=str, default=None)
|
||||
parser.add_argument("--model_path", type=str, default="LanguageBind/Open-Sora-Plan-v1.0.0")
|
||||
parser.add_argument("--version", type=str, default=None, choices=[None, "65x512x512", "221x512x512", "513x512x512"])
|
||||
parser.add_argument("--num_frames", type=int, default=1)
|
||||
parser.add_argument("--height", type=int, default=512)
|
||||
parser.add_argument("--width", type=int, default=512)
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
parser.add_argument("--ae", type=str, default="CausalVAEModel_4x8x8")
|
||||
parser.add_argument("--ae_path", type=str, default="CausalVAEModel_4x8x8")
|
||||
parser.add_argument("--text_encoder_name", type=str, default="DeepFloyd/t5-v1_1-xxl")
|
||||
parser.add_argument("--save_img_path", type=str, default="./sample_videos/t2v")
|
||||
parser.add_argument("--guidance_scale", type=float, default=7.5)
|
||||
parser.add_argument("--sample_method", type=str, default="PNDM")
|
||||
parser.add_argument("--num_sampling_steps", type=int, default=50)
|
||||
parser.add_argument("--fps", type=int, default=24)
|
||||
parser.add_argument("--run_time", type=int, default=0)
|
||||
parser.add_argument("--text_prompt", nargs="+")
|
||||
parser.add_argument("--force_images", action="store_true")
|
||||
parser.add_argument("--tile_overlap_factor", type=float, default=0.25)
|
||||
parser.add_argument("--enable_tiling", action="store_true")
|
||||
|
||||
# fvd
|
||||
parser.add_argument("--spatial_broadcast", action="store_true", help="Enable spatial attention skip")
|
||||
parser.add_argument(
|
||||
"--spatial_threshold", type=int, nargs=2, default=[100, 800], help="Spatial attention threshold"
|
||||
)
|
||||
parser.add_argument("--spatial_gap", type=int, default=2, help="Spatial attention gap")
|
||||
parser.add_argument("--temporal_broadcast", action="store_true", help="Enable temporal attention skip")
|
||||
parser.add_argument(
|
||||
"--temporal_threshold", type=int, nargs=2, default=[100, 800], help="Temporal attention threshold"
|
||||
)
|
||||
parser.add_argument("--temporal_gap", type=int, default=4, help="Temporal attention gap")
|
||||
parser.add_argument("--cross_broadcast", action="store_true", help="Enable cross attention skip")
|
||||
parser.add_argument("--cross_threshold", type=int, nargs=2, default=[100, 850], help="Cross attention threshold")
|
||||
parser.add_argument("--cross_gap", type=int, default=6, help="Cross attention gap")
|
||||
parser.add_argument(
|
||||
"--diffusion_skip",
|
||||
action="store_true",
|
||||
)
|
||||
parser.add_argument("--diffusion_skip_timestep", nargs="+")
|
||||
|
||||
args = parser.parse_args()
|
||||
config_args = OmegaConf.load(args.config)
|
||||
args = merge_args(args, config_args)
|
||||
|
||||
main(args)
|
||||
Executable
+1
@@ -0,0 +1 @@
|
||||
torchrun --standalone --nproc_per_node=1 scripts/opensora_plan/sample.py --config configs/opensora_plan/sample_65f.yaml
|
||||
Executable
+2
@@ -0,0 +1,2 @@
|
||||
# Pyramid Attention Broadcast
|
||||
torchrun --standalone --nproc_per_node=8 scripts/opensora_plan/sample.py --config configs/opensora_plan/sample_65f_pab.yaml
|
||||
Reference in New Issue
Block a user