Compare commits
12
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f008df2a00 | ||
|
|
01dbcfa163 | ||
|
|
30425102f2 | ||
|
|
818bb5484c | ||
|
|
62201e36f1 | ||
|
|
055e52e5ea | ||
|
|
7d2069596b | ||
|
|
c45009c9a4 | ||
|
|
b91020b407 | ||
|
|
2dcc5ea4f6 | ||
|
|
359151d9a0 | ||
|
|
ce67cd3729 |
+18
-18
@@ -117,11 +117,11 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
- "csrc/attn/config_vsa.py"
|
||||
- "csrc/attn/vsa.cpp"
|
||||
- "csrc/attn/video_sparse_attn/**"
|
||||
- "csrc/attn/video_sparse_attn/tk/**"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
- "csrc/attn/video_sparse_attn/config_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -133,10 +133,10 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
- "csrc/attn/st_attn.cpp"
|
||||
- "csrc/attn/sliding_tile_attn/**"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/sliding_tile_attn/config_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -147,10 +147,10 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
- "csrc/attn/st_attn.cpp"
|
||||
- "csrc/attn/sliding_tile_attn/**"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/sliding_tile_attn/config_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -161,12 +161,12 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/video_sparse_attn/**"
|
||||
- "csrc/attn/video_sparse_attn/tk/**"
|
||||
- "csrc/attn/tests/test_vsa.py"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
- "csrc/attn/config_vsa.py"
|
||||
- "csrc/attn/vsa.cpp"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
- "csrc/attn/video_sparse_attn/config_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
|
||||
@@ -104,16 +104,17 @@ jobs:
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
sta-kernel-paths: &sta-kernel-paths
|
||||
- 'csrc/attn/st_attn/**'
|
||||
- 'csrc/attn/setup_sta.py'
|
||||
- 'csrc/attn/config_sta.py'
|
||||
- 'csrc/attn/st_attn.cpp'
|
||||
- 'csrc/attn/sliding_tile_attn/**'
|
||||
- 'csrc/attn/sliding_tile_attn/tk/**'
|
||||
- 'csrc/attn/sliding_tile_attn/setup.py'
|
||||
- 'csrc/attn/sliding_tile_attn/config_sta.py'
|
||||
- 'csrc/attn/sliding_tile_attn/st_attn.cpp'
|
||||
vsa-kernel-paths: &vsa-kernel-paths
|
||||
- 'csrc/attn/vsa/**'
|
||||
- 'csrc/attn/tk/**'
|
||||
- 'csrc/attn/setup_vsa.py'
|
||||
- 'csrc/attn/config_vsa.py'
|
||||
- 'csrc/attn/vsa.cpp'
|
||||
- 'csrc/attn/video_sparse_attn/**'
|
||||
- 'csrc/attn/video_sparse_attn/tk/**'
|
||||
- 'csrc/attn/video_sparse_attn/setup.py'
|
||||
- 'csrc/attn/video_sparse_attn/config_vsa.py'
|
||||
- 'csrc/attn/video_sparse_attn/vsa.cpp'
|
||||
vsa-paths: &vsa-paths
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
|
||||
@@ -5,7 +5,7 @@ on:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
@@ -23,13 +23,13 @@ jobs:
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/attn
|
||||
cd csrc/attn/sliding_tile_attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
@@ -144,13 +144,13 @@ jobs:
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn # Move into the correct folder
|
||||
cd csrc/attn/sliding_tile_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py bdist_wheel --dist-dir=dist
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/attn
|
||||
cd csrc/attn/sliding_tile_attn
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
|
||||
@@ -165,7 +165,7 @@ jobs:
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/attn/dist/*.whl
|
||||
path: csrc/attn/sliding_tile_attn/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
@@ -239,11 +239,11 @@ jobs:
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn # Move into the correct folder
|
||||
cd csrc/attn/sliding_tile_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py sdist --dist-dir=dist
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/attn/dist/
|
||||
packages-dir: csrc/attn/sliding_tile_attn/dist/
|
||||
|
||||
@@ -5,7 +5,7 @@ on:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
@@ -23,13 +23,13 @@ jobs:
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/attn
|
||||
cd csrc/attn/video_sparse_attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_vsa.py)
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup_vsa.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
@@ -152,13 +152,13 @@ jobs:
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn # Move into the correct folder
|
||||
cd csrc/attn/video_sparse_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_vsa.py bdist_wheel --dist-dir=dist
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/attn
|
||||
cd csrc/attn/video_sparse_attn
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
|
||||
@@ -173,7 +173,7 @@ jobs:
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/attn/dist/*.whl
|
||||
path: csrc/attn/video_sparse_attn/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
@@ -247,11 +247,11 @@ jobs:
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn # Move into the correct folder
|
||||
cd csrc/attn/video_sparse_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_vsa.py sdist --dist-dir=dist
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/attn/dist/
|
||||
packages-dir: csrc/attn/video_sparse_attn/dist/
|
||||
|
||||
+6
-2
@@ -1,3 +1,7 @@
|
||||
[submodule "csrc/attn/tk"]
|
||||
path = csrc/attn/tk
|
||||
[submodule "csrc/attn/video_sparse_attn/tk"]
|
||||
path = csrc/attn/video_sparse_attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
[submodule "csrc/attn/sliding_tile_attn/tk"]
|
||||
path = csrc/attn/sliding_tile_attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
+2
-2
@@ -25,9 +25,9 @@ sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.4)
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
@@ -1,4 +0,0 @@
|
||||
off_hz = tl.program_id(2)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
@@ -1,2 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config.py
|
||||
include config_sta.py
|
||||
@@ -0,0 +1,87 @@
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We only support H100 for STA.
|
||||
|
||||
### Installation
|
||||
```bash
|
||||
pip install st_attn
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Usage
|
||||
End-2-end inference with FastVideo:
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
If you want to use sliding tile attention in your custom model:
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
```
|
||||
|
||||
|
||||
### Test
|
||||
```bash
|
||||
python ../tests/test_sta.py # test STA
|
||||
python ../tests/test_vsa.py # test VSA
|
||||
```
|
||||
### Benchmark
|
||||
```bash
|
||||
python ../benchmarks/bench_sta.py
|
||||
```
|
||||
|
||||
|
||||
### How Does STA Work?
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
|
||||
## Why is STA Fast?
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
|
||||
STA removes mixed blocks.
|
||||
|
||||
|
||||
<div align="center">
|
||||
<img src=../../../assets/sliding_tile_attn_map.png width="80%"/>
|
||||
</div>
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from csrc.attn.config_sta import kernels, sources, target
|
||||
from config_sta import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
@@ -9,7 +9,7 @@ target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "st_attn"
|
||||
VERSION = "0.0.4"
|
||||
VERSION = "0.0.6"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
|
||||
Submodule
+1
Submodule csrc/attn/sliding_tile_attn/tk added at 6c27e28c81
-1
Submodule csrc/attn/tk deleted from 1719fb7264
@@ -0,0 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config_vsa.py
|
||||
@@ -0,0 +1,61 @@
|
||||
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Video Sparse Attention (VSA)
|
||||
|
||||
### Installation
|
||||
We support H100 (via TK) and any other GPU (via triton) for VSA.
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
# test numerical
|
||||
python ../tests/test_vsa.py
|
||||
# (For H100) test speed
|
||||
python ../benchmarks/bench_vsa_hopper.py
|
||||
```
|
||||
|
||||
bench_vsa_hopper.py should print something like this:
|
||||
|
||||
```bash
|
||||
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
|
||||
|
||||
=== BLOCK SPARSE ATTENTION BENCHMARK ===
|
||||
Block Sparse Forward - TFLOPS: 5622.26
|
||||
Block Sparse Backward - TFLOPS: 3865.68
|
||||
```
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
@@ -9,10 +9,10 @@ target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "vsa"
|
||||
VERSION = "0.0.1"
|
||||
VERSION = "0.0.3"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn/video_sparse_attn"
|
||||
|
||||
# Set environment variables
|
||||
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
|
||||
Submodule
+1
Submodule csrc/attn/video_sparse_attn/tk added at 6c27e28c81
@@ -58,15 +58,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_sta.py install
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_vsa.py install
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
EXPOSE 22
|
||||
|
||||
@@ -58,15 +58,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_sta.py install
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_vsa.py install
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
EXPOSE 22
|
||||
|
||||
@@ -58,15 +58,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_sta.py install
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_vsa.py install
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
EXPOSE 22
|
||||
|
||||
@@ -58,15 +58,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_sta.py install
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_vsa.py install
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
EXPOSE 22
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
You can install the Sliding Tile Attention package using
|
||||
|
||||
```
|
||||
pip install st_attn==0.0.4
|
||||
pip install st_attn
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
@@ -12,7 +12,6 @@ We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have impleme
|
||||
First, install C++20 for ThunderKittens:
|
||||
|
||||
```bash
|
||||
cd csrc/sliding_tile_attention/
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
@@ -22,14 +21,20 @@ sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
|
||||
Install STA:
|
||||
Set up CUDA environment (if using CUDA 12.4):
|
||||
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
Install STA:
|
||||
|
||||
```bash
|
||||
cd csrc/attn/sliding_tile_attn/
|
||||
git submodule update --init --recursive
|
||||
python setup_sta.py install
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
# 🧪 Test
|
||||
|
||||
@@ -4,8 +4,7 @@
|
||||
You can install the Video Sparse Attention package using
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
@@ -34,9 +33,9 @@ export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
Install VSA:
|
||||
|
||||
```bash
|
||||
cd csrc/attn/
|
||||
cd csrc/attn/video_sparse_attn/
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
# 🧪 Test
|
||||
|
||||
+5
-15
@@ -39,7 +39,7 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=1
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation:
|
||||
# Model paths for DMD distillation:
|
||||
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
@@ -51,7 +51,7 @@ VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd
|
||||
--output_dir "checkpoints/SFwan_t2v_finetune"
|
||||
--max_train_steps 500
|
||||
--train_batch_size 1
|
||||
@@ -60,11 +60,11 @@ training_args=(
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
|
||||
--num_frames 81
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--log_visualization
|
||||
--simulate_generator_forward
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
--num_frame_per_block 3
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
)
|
||||
@@ -138,15 +138,6 @@ dmd_args=(
|
||||
--fake_score_betas '0.0,0.999'
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks False # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
--validate_cache_structure False # Set to True for debugging KV cache issues
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
@@ -161,5 +152,4 @@ srun torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
@@ -4,9 +4,7 @@ These are end-to-end example scripts for distilling Wan2.1 T2V 1.3B model using
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
cd csrc/attn
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
|
||||
@@ -4,9 +4,7 @@ These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
cd csrc/attn
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
### Data-free Distillation
|
||||
|
||||
@@ -4,9 +4,7 @@ These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
cd csrc/attn
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
|
||||
@@ -28,4 +28,4 @@ def main():
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_5B_ti2v"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model_name = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
# I2V is triggered just by passing in an image_path argument
|
||||
prompt = "Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside."
|
||||
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, image_path=image_path)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
|
||||
# T2V mode
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -7,9 +7,7 @@ These are e2e example scripts for finetuning Wan2.1 T2V with VSA to accelerate i
|
||||
## Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
cd csrc/attn
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
### Download the synthetic dataset:
|
||||
|
||||
@@ -85,6 +85,9 @@ class PipelineConfig:
|
||||
# DMD parameters
|
||||
dmd_denoising_steps: list[int] | None = field(default=None)
|
||||
|
||||
# Wan2.2 TI2V parameters
|
||||
ti2v_task: bool = False
|
||||
|
||||
# Compilation
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ from fastvideo.configs.pipelines.wan import (FastWan2_1_T2V_480P_Config,
|
||||
SelfForcingWanT2V480PConfig,
|
||||
WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
verify_model_config_and_directory)
|
||||
|
||||
@@ -57,7 +57,7 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@@ -148,4 +148,4 @@ class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
|
||||
is_causal: bool = True
|
||||
flow_shift: int = 5
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 750, 500, 250])
|
||||
default_factory=lambda: [1000, 750, 500, 250])
|
||||
|
||||
@@ -10,6 +10,7 @@ from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
# isort: off
|
||||
from fastvideo.configs.sample.wan import (
|
||||
FastWanT2V480PConfig,
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
Wan2_2_I2V_A14B_SamplingParam,
|
||||
Wan2_2_T2V_A14B_SamplingParam,
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
|
||||
@@ -107,6 +107,21 @@ class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
|
||||
fps: int = 16
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.1 Fun Models =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
guidance_scale: float = 6.0
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.2 TI2V Models =============
|
||||
# =============================================
|
||||
@@ -148,4 +163,4 @@ class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
# =============================================
|
||||
@dataclass
|
||||
class SelfForcingWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
|
||||
pass
|
||||
pass
|
||||
|
||||
@@ -119,7 +119,7 @@ class FastVideoArgs:
|
||||
output_type: str = "pil"
|
||||
|
||||
# CPU offload parameters
|
||||
dit_cpu_offload: bool = True
|
||||
dit_cpu_offload: bool = False
|
||||
use_fsdp_inference: bool = True
|
||||
text_encoder_cpu_offload: bool = True
|
||||
image_encoder_cpu_offload: bool = True
|
||||
@@ -691,9 +691,6 @@ class TrainingArgs(FastVideoArgs):
|
||||
enable_gradient_masking: bool = True
|
||||
gradient_mask_last_n_frames: int = 21
|
||||
validate_cache_structure: bool = False # Debug flag for cache validation
|
||||
same_step_across_blocks: bool = False # Use same exit timestep for all blocks
|
||||
last_step_only: bool = False # Only use the last timestep for training
|
||||
context_noise: int = 0 # Context noise level for cache updates
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
@@ -1102,19 +1099,6 @@ class TrainingArgs(FastVideoArgs):
|
||||
"--validate-cache-structure",
|
||||
action=StoreBoolean,
|
||||
help="Whether to validate KV cache structure (debug flag)")
|
||||
parser.add_argument(
|
||||
"--same-step-across-blocks",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use the same exit timestep for all blocks")
|
||||
parser.add_argument(
|
||||
"--last-step-only",
|
||||
action=StoreBoolean,
|
||||
help="Whether to only use the last timestep for training")
|
||||
parser.add_argument(
|
||||
"--context-noise",
|
||||
type=int,
|
||||
default=TrainingArgs.context_noise,
|
||||
help="Context noise level for cache updates")
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
@@ -9,9 +9,6 @@ import torch.nn.functional as F
|
||||
from fastvideo.layers.custom_op import CustomOp
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@CustomOp.register("rms_norm")
|
||||
class RMSNorm(CustomOp):
|
||||
@@ -103,13 +100,7 @@ class ScaleResidual(nn.Module):
|
||||
def forward(self, residual: torch.Tensor, x: torch.Tensor,
|
||||
gate: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply gated residual connection."""
|
||||
# logger.info("x.shape: %s", x.shape)
|
||||
# if isinstance(gate, torch.Tensor):
|
||||
# logger.info("gate.shape: %s", gate.shape)
|
||||
|
||||
num_frames = gate.shape[1]
|
||||
frame_seqlen = x.shape[1] // num_frames
|
||||
return residual + (x.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * gate).flatten(1, 2)
|
||||
return residual + x * gate
|
||||
|
||||
|
||||
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
|
||||
@@ -181,35 +172,11 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
but before normalization)
|
||||
"""
|
||||
# Apply residual connection with gating
|
||||
# logger.info("x.shape: %s", x.shape)
|
||||
if isinstance(gate, int):
|
||||
# used by cross-attention, should be 1
|
||||
assert gate == 1
|
||||
residual_output = residual + x * gate
|
||||
elif isinstance(gate, torch.Tensor):
|
||||
# logger.info("gate.shape: %s", gate.shape)
|
||||
if gate.dim() == 3:
|
||||
# used by bidirectional self attention
|
||||
residual_output = residual + x * gate
|
||||
else:
|
||||
assert gate.dim() == 4
|
||||
num_frames = gate.shape[1]
|
||||
frame_seqlen = x.shape[1] // num_frames
|
||||
residual_output = residual + (x.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * gate).flatten(1, 2)
|
||||
# residual_output = residual + x * gate
|
||||
else:
|
||||
raise ValueError(f"Gate type {type(gate)} not supported")
|
||||
# logger.info("residual_output.shape: %s", residual_output.shape)
|
||||
|
||||
residual_output = residual + x * gate
|
||||
# Apply normalization
|
||||
normalized = self.norm(residual_output)
|
||||
# Apply scale and shift
|
||||
if isinstance(scale, torch.Tensor) and scale.dim() == 4:
|
||||
num_frames = scale.shape[1]
|
||||
frame_seqlen = normalized.shape[1] // num_frames
|
||||
modulated = (normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1.0 + scale) + shift).flatten(1, 2)
|
||||
else:
|
||||
modulated = normalized * (1.0 + scale) + shift
|
||||
modulated = normalized * (1.0 + scale) + shift
|
||||
return modulated, residual_output
|
||||
|
||||
|
||||
@@ -252,15 +219,7 @@ class LayerNormScaleShift(nn.Module):
|
||||
scale: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply ln followed by scale and shift in a single fused operation."""
|
||||
normalized = self.norm(x)
|
||||
if scale.dim() == 4:
|
||||
num_frames = scale.shape[1]
|
||||
frame_seqlen = normalized.shape[1] // num_frames
|
||||
if self.compute_dtype == torch.float32:
|
||||
return (normalized.float().unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1.0 + scale) + shift).flatten(1, 2).to(x.dtype)
|
||||
else:
|
||||
return (normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1.0 + scale) + shift).flatten(1, 2)
|
||||
if self.compute_dtype == torch.float32:
|
||||
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
|
||||
else:
|
||||
if self.compute_dtype == torch.float32:
|
||||
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
|
||||
else:
|
||||
return normalized * (1.0 + scale) + shift
|
||||
return normalized * (1.0 + scale) + shift
|
||||
|
||||
@@ -33,8 +33,6 @@ from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
def _rotate_neox(x: torch.Tensor) -> torch.Tensor:
|
||||
x1 = x[..., :x.shape[-1] // 2]
|
||||
x2 = x[..., x.shape[-1] // 2:]
|
||||
|
||||
@@ -176,4 +176,4 @@ def unpatchify(x, t, h, w, patch_size, channels) -> torch.Tensor:
|
||||
x = torch.einsum("nthwcopq->nctohpwq", x)
|
||||
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
|
||||
|
||||
return imgs
|
||||
return imgs
|
||||
|
||||
@@ -147,10 +147,6 @@ class CausalWanSelfAttention(nn.Module):
|
||||
# Assign new keys/values directly up to current_end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
|
||||
kv_cache["k"] = kv_cache["k"].clone()
|
||||
kv_cache["v"] = kv_cache["v"].clone()
|
||||
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
x = self.attn(
|
||||
@@ -248,36 +244,19 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
current_start: int = 0,
|
||||
cache_start: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
logger.info("temb.shape: %s", temb.shape)
|
||||
num_frames = temb.shape[1]
|
||||
logger.info("first hidden_states.shape: %s", hidden_states.shape)
|
||||
logger.info("num_frames: %s", num_frames)
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
frame_seqlen = hidden_states.shape[1] // temb.shape[1]
|
||||
logger.info("frame_seqlen: %s", frame_seqlen)
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb.float()
|
||||
logger.info("e.shape: %s", e.shape)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=2)
|
||||
6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
logger.info("hidden_states.shape: %s", hidden_states.shape)
|
||||
logger.info("scale_msa.shape: %s", scale_msa.shape)
|
||||
logger.info("shift_msa.shape: %s", shift_msa.shape)
|
||||
|
||||
norm_hidden_states_unflattened = self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
|
||||
# logger.info("norm_hidden_states_unflattened.shape: %s", norm_hidden_states_unflattened.shape)
|
||||
|
||||
# norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
# (1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
norm_hidden_states = (norm_hidden_states_unflattened *
|
||||
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
|
||||
# logger.info("1 norm_hidden_states.shape: %s", norm_hidden_states.shape)
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
@@ -299,8 +278,6 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
# logger.info("after self_attn_residual_norm norm_hidden_states.shape: %s", norm_hidden_states.shape)
|
||||
# logger.info("after self_attn_residual_norm hidden_states.shape: %s", hidden_states.shape)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
@@ -311,16 +288,12 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
crossattn_cache=crossattn_cache)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
# logger.info("after cross_attn_residual_norm norm_hidden_states.shape: %s", norm_hidden_states.shape)
|
||||
# logger.info("after cross_attn_residual_norm hidden_states.shape: %s", hidden_states.shape)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
# logger.info("after mlp_residual norm_hidden_states.shape: %s", norm_hidden_states.shape)
|
||||
logger.info("after mlp_residual hidden_states.shape: %s", hidden_states.shape)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
@@ -386,10 +359,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
# Debug: Log configuration values
|
||||
proj_out_dim = config.out_channels * math.prod(config.patch_size)
|
||||
|
||||
self.proj_out = nn.Linear(inner_dim, proj_out_dim)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
|
||||
|
||||
@@ -478,7 +449,6 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
This function will be run for num_frame times.
|
||||
Process the latent frames one by one (1560 tokens each)
|
||||
"""
|
||||
# logger.info("forward inference hidden_states.shape: %s", hidden_states.shape)
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
@@ -515,13 +485,10 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
# logger.info("forward inference flattened and transposed hidden_states.shape: %s", hidden_states.shape)
|
||||
|
||||
logger.info("timestep shape: %s", timestep.shape)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states = torch.concat(
|
||||
@@ -559,14 +526,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
**causal_kwargs)
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
# logger.info("===== INFERENCE 5. Output norm, projection & unpatchify")
|
||||
# logger.info("hidden_states.shape: %s", hidden_states.shape)
|
||||
# logger.info("temb.shape: %s", temb.shape)
|
||||
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
|
||||
# logger.info("WTFWTF train temb.shape: %s", temb.shape)
|
||||
# logger.info("WTFWTF train self.scale_shift_table.shape: %s", self.scale_shift_table.shape)
|
||||
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
|
||||
dim=2)
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
|
||||
dim=1)
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
@@ -588,8 +549,6 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
start_frame: int = 0,
|
||||
**kwargs) -> torch.Tensor:
|
||||
|
||||
# logger.info("===== forward train hidden_states.shape: %s", hidden_states.shape)
|
||||
# logger.info("===== forward train timestep.shape: %s", timestep.shape)
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
@@ -635,14 +594,10 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
# logger.info("forward train flattened and transposed hidden_states.shape: %s", hidden_states.shape)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
# logger.info("forward train timestep_proj.shape: %s", timestep_proj.shape)
|
||||
# logger.info("forward train timestep.shape: %s", timestep.shape)
|
||||
# logger.info("forward train temb.shape: %s", temb.shape)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states = torch.concat(
|
||||
@@ -662,35 +617,16 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
timestep_proj, freqs_cis,
|
||||
block_mask=self.block_mask)
|
||||
else:
|
||||
for block_index, block in enumerate(self.blocks):
|
||||
logger.info("===== TRAIN block %d", block_index)
|
||||
logger.info("hidden_states.shape: %s", hidden_states.shape)
|
||||
# logger.info("encoder_hidden_states.shape: %s", encoder_hidden_states.shape)
|
||||
logger.info("timestep_proj.shape: %s", timestep_proj.shape)
|
||||
# logger.info("freqs_cis.shape: %s", freqs_cis.shape)
|
||||
# logger.info("block_mask.shape: %s", self.block_mask.shape)
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis,
|
||||
block_mask=self.block_mask)
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
# logger.info("===== TRAIN 5. Output norm, projection & unpatchify")
|
||||
# logger.info("hidden_states.shape: %s", hidden_states.shape)
|
||||
# logger.info("temb.shape: %s", temb.shape)
|
||||
# shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
|
||||
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
|
||||
# logger.info("WTFWTF train temb.shape: %s", temb.shape)
|
||||
# logger.info("WTFWTF train self.scale_shift_table.shape: %s", self.scale_shift_table.shape)
|
||||
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
|
||||
dim=2)
|
||||
# logger.info("DEBUG scale.shape: %s", scale.shape)
|
||||
# logger.info("DEBUG shift.shape: %s", shift.shape)
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
|
||||
dim=1)
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
# logger.info("DEBUG after proj_out hidden_states.shape: %s", hidden_states.shape)
|
||||
# logger.info(f"DEBUG reshape dimensions: batch_size={batch_size}, post_patch_num_frames={post_patch_num_frames}")
|
||||
# logger.info(f"DEBUG reshape dimensions: post_patch_height={post_patch_height}, post_patch_width={post_patch_width}")
|
||||
# logger.info(f"DEBUG patch dimensions: p_t={p_t}, p_h={p_h}, p_w={p_w}")
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
@@ -709,4 +645,4 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
if kwargs.get('kv_cache', None) is not None:
|
||||
return self._forward_inference(*args, **kwargs)
|
||||
else:
|
||||
return self._forward_train(*args, **kwargs)
|
||||
return self._forward_train(*args, **kwargs)
|
||||
|
||||
@@ -636,14 +636,8 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
timestep: torch.IntTensor,
|
||||
) -> torch.Tensor:
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
# TODO: hack
|
||||
if timestep.ndim == 2 and timestep.shape[1] == 1:
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
logger.info("self.timesteps shape: %s", self.timesteps.shape)
|
||||
logger.info("timestep shape: %s", timestep.shape)
|
||||
if timestep.ndim > 1:
|
||||
timestep = timestep.squeeze(0)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
@@ -6,9 +6,6 @@ from typing import Any
|
||||
import torch
|
||||
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# TODO(PY): move it elsewhere
|
||||
def auto_attributes(init_func):
|
||||
"""
|
||||
@@ -149,9 +146,6 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
|
||||
"""
|
||||
Convert predicted noise to clean latent.
|
||||
"""
|
||||
logger.info(f"timestep: {timestep.shape}")
|
||||
logger.info(f"noise_input_latent: {noise_input_latent.shape}")
|
||||
logger.info(f"pred_noise: {pred_noise.shape}")
|
||||
timestep = timestep.expand(noise_input_latent.shape[0])
|
||||
dtype = pred_noise.dtype
|
||||
device = pred_noise.device
|
||||
@@ -163,4 +157,4 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
pred_video = noise_input_latent - sigma_t * pred_noise
|
||||
return pred_video.to(dtype)
|
||||
return pred_video.to(dtype)
|
||||
|
||||
@@ -27,7 +27,6 @@ from fastvideo.layers.activation import get_act_fn
|
||||
from fastvideo.models.vaes.common import (DiagonalGaussianDistribution,
|
||||
ParallelTiledVAE)
|
||||
from fastvideo.platforms import current_platform
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
CACHE_T = 2
|
||||
|
||||
@@ -36,7 +35,6 @@ feat_cache = contextvars.ContextVar("feat_cache", default=None)
|
||||
feat_idx = contextvars.ContextVar("feat_idx", default=0)
|
||||
first_chunk = contextvars.ContextVar("first_chunk", default=None)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@contextmanager
|
||||
def forward_context(first_frame_arg=False,
|
||||
@@ -1131,7 +1129,6 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
self._conv_idx = 0
|
||||
self._feat_map = [None] * self._conv_num
|
||||
# cache encode
|
||||
logger.info("self.config.load_encoder: %s", self.config.load_encoder)
|
||||
if self.config.load_encoder:
|
||||
self._enc_conv_num = _count_conv3d(self.encoder)
|
||||
self._enc_conv_idx = 0
|
||||
|
||||
@@ -63,6 +63,7 @@ class WanPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
|
||||
@@ -228,10 +228,7 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
dim=2)
|
||||
|
||||
# Prepare inputs
|
||||
t_expand = t_cur.expand(latent_model_input.shape[0])
|
||||
# t_expand = t_cur * torch.ones((latent_model_input.shape[0], 1), device=latent_model_input.device, dtype=torch.long)
|
||||
# t_expand = t_expand.repeat(1, self.sliding_window_num_frames)
|
||||
|
||||
t_expand = t_cur.repeat(latent_model_input.shape[0])
|
||||
|
||||
# Attention metadata if needed
|
||||
if (vsa_available and self.attn_backend
|
||||
@@ -265,11 +262,10 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
# Run transformer; follow DMD stage pattern
|
||||
t_expanded_noise= t_cur * torch.ones((latent_model_input.shape[0], 1), device=latent_model_input.device, dtype=torch.long)
|
||||
pred_noise_btchw = self.transformer(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
t_expanded_noise,
|
||||
t_expand,
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
@@ -330,11 +326,10 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
set_forward_context(current_timestep=0,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
t_expanded_context = t_context * torch.ones((context_bcthw.shape[0], 1), device=context_bcthw.device, dtype=torch.long)
|
||||
_ = self.transformer(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
t_expanded_context,
|
||||
t_context,
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
@@ -411,4 +406,4 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
"is_init":
|
||||
False,
|
||||
})
|
||||
self.crossattn_cache = crossattn_cache
|
||||
self.crossattn_cache = crossattn_cache
|
||||
|
||||
@@ -4,6 +4,7 @@ Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
import math
|
||||
import weakref
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
@@ -24,12 +25,13 @@ from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.utils import dict_to_3d_list
|
||||
from fastvideo.utils import dict_to_3d_list, masks_like
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
@@ -60,11 +62,13 @@ class DenoisingStage(PipelineStage):
|
||||
transformer,
|
||||
scheduler,
|
||||
pipeline=None,
|
||||
transformer_2=None) -> None:
|
||||
transformer_2=None,
|
||||
vae=None) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.transformer_2 = transformer_2
|
||||
self.scheduler = scheduler
|
||||
self.vae = vae
|
||||
self.pipeline = weakref.ref(pipeline) if pipeline else None
|
||||
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
|
||||
self.attn_backend = get_attn_backend(
|
||||
@@ -193,6 +197,44 @@ class DenoisingStage(PipelineStage):
|
||||
boundary_timestep = fastvideo_args.boundary_ratio * self.scheduler.num_train_timesteps
|
||||
else:
|
||||
boundary_timestep = None
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
assert latent_model_input.shape[0] == 1, "only support batch size 1"
|
||||
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
# TI2V directly replaces the first frame of the latent with
|
||||
# the image latent instead of appending along the channel dim
|
||||
assert batch.image_latent is None, "TI2V task should not have image latents"
|
||||
assert self.vae is not None, "VAE is not provided for TI2V task"
|
||||
z = self.vae.encode(batch.pil_image).mean.float()
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
z -= self.vae.shift_factor.to(z.device, z.dtype)
|
||||
else:
|
||||
z -= self.vae.shift_factor
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
z = z * self.vae.scaling_factor.to(z.device, z.dtype)
|
||||
else:
|
||||
z = z * self.vae.scaling_factor
|
||||
|
||||
latent_model_input = latent_model_input.squeeze(0)
|
||||
_, mask2 = masks_like([latent_model_input], zero=True)
|
||||
|
||||
latent_model_input = (1. -
|
||||
mask2[0]) * z + mask2[0] * latent_model_input
|
||||
# latent_model_input = latent_model_input.unsqueeze(0)
|
||||
latent_model_input = latent_model_input.to(get_local_torch_device())
|
||||
latents = latent_model_input
|
||||
F = batch.num_frames
|
||||
temporal_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_temporal
|
||||
spatial_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
seq_len = ((F - 1) // temporal_scale +
|
||||
1) * (batch.height // spatial_scale) * (
|
||||
batch.width // spatial_scale) // (patch_size[1] *
|
||||
patch_size[2])
|
||||
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
@@ -217,19 +259,32 @@ class DenoisingStage(PipelineStage):
|
||||
self.transformer.to('cpu')
|
||||
current_model = self.transformer_2
|
||||
current_guidance_scale = batch.guidance_scale_2
|
||||
assert current_model is not None, "current_model is None"
|
||||
|
||||
# Expand latents for I2V
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if batch.image_latent is not None:
|
||||
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.image_latent],
|
||||
dim=1).to(target_dtype)
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
timestep = torch.stack([t]).to(get_local_torch_device())
|
||||
temp_ts = (mask2[0][0][:, ::2, ::2] * timestep).flatten()
|
||||
temp_ts = torch.cat([
|
||||
temp_ts,
|
||||
temp_ts.new_ones(seq_len - temp_ts.size(0)) * timestep
|
||||
])
|
||||
timestep = temp_ts.unsqueeze(0)
|
||||
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
|
||||
else:
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
|
||||
assert torch.isnan(latent_model_input).sum() == 0
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
|
||||
# Prepare inputs for transformer
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (
|
||||
torch.tensor(
|
||||
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
|
||||
@@ -329,6 +384,11 @@ class DenoisingStage(PipelineStage):
|
||||
latents,
|
||||
**extra_step_kwargs,
|
||||
return_dict=False)[0]
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
latents = latents.squeeze(0)
|
||||
latents = (1. - mask2[0]) * z + mask2[0] * latents
|
||||
# latents = latents.unsqueeze(0)
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
(i + 1) > num_warmup_steps and
|
||||
|
||||
@@ -4,6 +4,8 @@ Input validation stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torchvision.transforms.functional as TF
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -12,6 +14,7 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import (StageValidators,
|
||||
VerificationResult)
|
||||
from fastvideo.utils import best_output_size
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -94,6 +97,7 @@ class InputValidationStage(PipelineStage):
|
||||
)
|
||||
|
||||
# for i2v, get image from image_path
|
||||
# @TODO(Wei) hard-coded for wan2.2 5b ti2v for now. Should put this in image_encoding stage
|
||||
if batch.image_path is not None:
|
||||
if batch.image_path.endswith(".mp4"):
|
||||
image = load_video(batch.image_path)[0]
|
||||
@@ -101,6 +105,36 @@ class InputValidationStage(PipelineStage):
|
||||
image = load_image(batch.image_path)
|
||||
batch.pil_image = image
|
||||
|
||||
# further processing for ti2v task
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
img = batch.pil_image
|
||||
ih, iw = img.height, img.width
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
|
||||
max_area = 704 * 1280
|
||||
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
|
||||
|
||||
scale = max(ow / iw, oh / ih)
|
||||
img = img.resize((round(iw * scale), round(ih * scale)),
|
||||
Image.LANCZOS)
|
||||
logger.info("resized img height: %s, img width: %s", img.height,
|
||||
img.width)
|
||||
|
||||
# center-crop
|
||||
x1 = (img.width - ow) // 2
|
||||
y1 = (img.height - oh) // 2
|
||||
img = img.crop((x1, y1, x1 + ow, y1 + oh))
|
||||
assert img.width == ow and img.height == oh
|
||||
|
||||
# to tensor
|
||||
img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(
|
||||
self.device).unsqueeze(1)
|
||||
img = img.unsqueeze(0)
|
||||
batch.height = oh
|
||||
batch.width = ow
|
||||
batch.pil_image = img
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
|
||||
@@ -52,40 +52,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
super().initialize_training_pipeline(training_args)
|
||||
self.dfake_gen_update_ratio = getattr(training_args, 'dfake_gen_update_ratio', 5)
|
||||
|
||||
self.num_frame_per_block = getattr(training_args, 'num_frame_per_block', 3)
|
||||
self.independent_first_frame = getattr(training_args, 'independent_first_frame', False)
|
||||
self.same_step_across_blocks = getattr(training_args, 'same_step_across_blocks', False)
|
||||
self.last_step_only = getattr(training_args, 'last_step_only', False)
|
||||
self.context_noise = getattr(training_args, 'context_noise', 0)
|
||||
|
||||
self.frame_seq_length = 1560 # TODO: Calculate this dynamically based on patch size
|
||||
|
||||
self.kv_cache1 = None
|
||||
self.crossattn_cache = None
|
||||
|
||||
logger.info(f"Self-forcing generator update ratio: {self.dfake_gen_update_ratio}")
|
||||
|
||||
def generate_and_sync_list(self, num_blocks, num_denoising_steps, device):
|
||||
"""Generate and synchronize random exit flags across distributed processes."""
|
||||
rank = dist.get_rank() if dist.is_initialized() else 0
|
||||
|
||||
if rank == 0:
|
||||
# Generate random indices
|
||||
indices = torch.randint(
|
||||
low=0,
|
||||
high=num_denoising_steps,
|
||||
size=(num_blocks,),
|
||||
device=device
|
||||
)
|
||||
if self.last_step_only:
|
||||
indices = torch.ones_like(indices) * (num_denoising_steps - 1)
|
||||
else:
|
||||
indices = torch.empty(num_blocks, dtype=torch.long, device=device)
|
||||
|
||||
if dist.is_initialized():
|
||||
dist.broadcast(indices, src=0) # Broadcast the random indices to all ranks
|
||||
return indices.tolist()
|
||||
|
||||
def generator_loss(self, training_batch: TrainingBatch) -> tuple[torch.Tensor, dict[str, Any]]:
|
||||
"""
|
||||
Compute generator loss using DMD-style approach.
|
||||
@@ -194,286 +162,223 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
return pred_video
|
||||
|
||||
def _generator_multi_step_simulation_forward(
|
||||
self, training_batch: TrainingBatch, return_sim_steps: bool = False) -> torch.Tensor:
|
||||
"""Forward pass through student transformer matching inference procedure with KV cache management.
|
||||
|
||||
This function is adapted from the reference self-forcing implementation's inference_with_trajectory
|
||||
and includes gradient masking logic for dynamic frame generation.
|
||||
"""
|
||||
self, training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""Forward pass through student transformer matching inference procedure with KV cache management."""
|
||||
latents = training_batch.latents
|
||||
dtype = latents.dtype
|
||||
batch_size = latents.shape[0]
|
||||
initial_latent = getattr(training_batch, 'image_latent', None)
|
||||
|
||||
# Dynamic frame generation logic (adapted from _run_generator)
|
||||
num_training_frames = getattr(self.training_args, 'num_frames', 21)
|
||||
|
||||
# During training, the number of generated frames should be uniformly sampled from
|
||||
# [21, self.num_training_frames], but still being a multiple of self.num_frame_per_block
|
||||
min_num_frames = 20 if self.independent_first_frame else 21
|
||||
max_num_frames = num_training_frames - 1 if self.independent_first_frame else num_training_frames
|
||||
assert max_num_frames % self.num_frame_per_block == 0
|
||||
assert min_num_frames % self.num_frame_per_block == 0
|
||||
max_num_blocks = max_num_frames // self.num_frame_per_block
|
||||
min_num_blocks = min_num_frames // self.num_frame_per_block
|
||||
|
||||
# Sample number of blocks and sync across processes
|
||||
num_generated_blocks = torch.randint(min_num_blocks, max_num_blocks + 1, (1,), device=self.device)
|
||||
if dist.is_initialized():
|
||||
dist.broadcast(num_generated_blocks, src=0)
|
||||
num_generated_blocks = num_generated_blocks.item()
|
||||
num_generated_frames = num_generated_blocks * self.num_frame_per_block
|
||||
if self.independent_first_frame and initial_latent is None:
|
||||
num_generated_frames += 1
|
||||
min_num_frames += 1
|
||||
|
||||
# Create noise with dynamic shape
|
||||
if initial_latent is not None:
|
||||
noise_shape = [batch_size, num_generated_frames - 1, *self.video_latent_shape[2:]]
|
||||
else:
|
||||
noise_shape = [batch_size, num_generated_frames, *self.video_latent_shape[2:]]
|
||||
|
||||
noise = torch.randn(noise_shape, device=self.device, dtype=dtype)
|
||||
if self.sp_world_size > 1:
|
||||
noise = rearrange(noise, "b (n t) c h w -> b n t c h w", n=self.sp_world_size).contiguous()
|
||||
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
|
||||
|
||||
batch_size, num_frames, num_channels, height, width = noise.shape
|
||||
|
||||
# Block size calculation
|
||||
if not self.independent_first_frame or (self.independent_first_frame and initial_latent is not None):
|
||||
assert num_frames % self.num_frame_per_block == 0
|
||||
num_blocks = num_frames // self.num_frame_per_block
|
||||
else:
|
||||
assert (num_frames - 1) % self.num_frame_per_block == 0
|
||||
num_blocks = (num_frames - 1) // self.num_frame_per_block
|
||||
|
||||
num_input_frames = initial_latent.shape[1] if initial_latent is not None else 0
|
||||
num_output_frames = num_frames + num_input_frames
|
||||
output = torch.zeros([batch_size, num_output_frames, num_channels, height, width],
|
||||
device=noise.device, dtype=noise.dtype)
|
||||
|
||||
# Step 1: Initialize KV cache to all zeros
|
||||
self.kv_cache1, self.crossattn_cache = self._initialize_simulation_caches(batch_size, dtype, self.device)
|
||||
# Step 1: Randomly sample a target timestep index from denoising_step_list
|
||||
target_timestep_idx = torch.randint(0,
|
||||
len(self.denoising_step_list), [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
target_timestep_idx_int = target_timestep_idx.item()
|
||||
target_timestep = self.denoising_step_list[target_timestep_idx]
|
||||
|
||||
# Step 2: Initialize KV cache and cross-attention cache for causal generation
|
||||
kv_cache, crossattn_cache = self._initialize_simulation_caches(batch_size, dtype, self.device)
|
||||
|
||||
# Validate cache structure (can be disabled in production)
|
||||
if getattr(self.training_args, 'validate_cache_structure', False):
|
||||
self._validate_cache_structure(self.kv_cache1, self.crossattn_cache, batch_size)
|
||||
self._validate_cache_structure(kv_cache, crossattn_cache, batch_size)
|
||||
|
||||
# Step 2: Cache context feature
|
||||
current_start_frame = 0
|
||||
if initial_latent is not None:
|
||||
timestep = torch.ones([batch_size, 1], device=noise.device, dtype=torch.int64) * 0
|
||||
output[:, :1] = initial_latent
|
||||
with torch.no_grad():
|
||||
# Build input kwargs for initial latent
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
initial_latent, timestep * 0, training_batch.conditional_dict, training_batch)
|
||||
# Step 3: Simulate the multi-step inference process up to the target timestep
|
||||
# Start from pure noise like in inference
|
||||
current_noise_latents = torch.randn(self.video_latent_shape,
|
||||
device=self.device,
|
||||
dtype=dtype)
|
||||
if self.sp_world_size > 1:
|
||||
current_noise_latents = rearrange(
|
||||
current_noise_latents,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
current_noise_latents = current_noise_latents[:, self.
|
||||
rank_in_sp_group, :, :, :, :]
|
||||
current_noise_latents_copy = current_noise_latents.clone()
|
||||
|
||||
# Step 4: Determine gradient masking for frame-level selective training
|
||||
num_frames = current_noise_latents.shape[1]
|
||||
num_frame_per_block = getattr(self.training_args, 'num_frame_per_block', 3)
|
||||
independent_first_frame = getattr(self.training_args, 'independent_first_frame', False)
|
||||
enable_gradient_masking = getattr(self.training_args, 'enable_gradient_masking', True)
|
||||
gradient_mask_last_n_frames = getattr(self.training_args, 'gradient_mask_last_n_frames', 21)
|
||||
|
||||
gradient_mask = None
|
||||
if enable_gradient_masking:
|
||||
# Calculate which frames should have gradients enabled (last N frames)
|
||||
start_gradient_frame_index = max(0, num_frames - gradient_mask_last_n_frames)
|
||||
gradient_mask = torch.ones(batch_size, num_frames, dtype=torch.bool, device=self.device)
|
||||
|
||||
self.transformer(
|
||||
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
|
||||
timestep=training_batch_temp.input_kwargs['timestep'],
|
||||
encoder_hidden_states_image=training_batch_temp.input_kwargs.get('encoder_hidden_states_image'),
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=current_start_frame * self.frame_seq_length
|
||||
if independent_first_frame:
|
||||
# First frame is independent, disable gradients for early frames
|
||||
gradient_mask[:, :max(1, start_gradient_frame_index)] = False
|
||||
else:
|
||||
# Disable gradients for early blocks
|
||||
num_early_blocks = start_gradient_frame_index // num_frame_per_block
|
||||
gradient_mask[:, :num_early_blocks * num_frame_per_block] = False
|
||||
|
||||
# Step 5: Multi-step simulation with causal block processing
|
||||
num_frames_per_block = getattr(self.training_args, 'num_frame_per_block', 3)
|
||||
|
||||
t = num_frames
|
||||
if not independent_first_frame or (independent_first_frame and hasattr(training_batch, 'image_latent') and training_batch.image_latent is not None):
|
||||
if t % num_frames_per_block != 0:
|
||||
raise ValueError(
|
||||
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
|
||||
)
|
||||
current_start_frame += 1
|
||||
num_blocks = t // num_frames_per_block
|
||||
block_sizes = [num_frames_per_block] * num_blocks
|
||||
else:
|
||||
if (t - 1) % num_frames_per_block != 0:
|
||||
raise ValueError(
|
||||
"(num_frames - 1) must be divisible by num_frame_per_block when independent_first_frame=True"
|
||||
)
|
||||
num_blocks = (t - 1) // num_frames_per_block
|
||||
block_sizes = [1] + [num_frames_per_block] * num_blocks
|
||||
|
||||
# Step 3: Temporal denoising loop
|
||||
all_num_frames = [self.num_frame_per_block] * num_blocks
|
||||
if self.independent_first_frame and initial_latent is None:
|
||||
all_num_frames = [1] + all_num_frames
|
||||
num_denoising_steps = len(self.denoising_step_list)
|
||||
exit_flags = self.generate_and_sync_list(len(all_num_frames), num_denoising_steps, device=noise.device)
|
||||
start_gradient_frame_index = num_output_frames - 21
|
||||
|
||||
for block_index, current_num_frames in enumerate(all_num_frames):
|
||||
noisy_input = noise[:, current_start_frame - num_input_frames:current_start_frame + current_num_frames - num_input_frames]
|
||||
|
||||
# Step 3.1: Spatial denoising loop
|
||||
for index, current_timestep in enumerate(self.denoising_step_list):
|
||||
if self.same_step_across_blocks:
|
||||
exit_flag = (index == exit_flags[0])
|
||||
else:
|
||||
exit_flag = (index == exit_flags[block_index])
|
||||
max_target_idx = len(self.denoising_step_list) - 1
|
||||
noise_latents = []
|
||||
noise_latent_index = target_timestep_idx_int - 1
|
||||
|
||||
if max_target_idx > 0:
|
||||
# Run student model for all steps before the target timestep with causal blocks
|
||||
with torch.no_grad():
|
||||
start_index = 0
|
||||
current_start_frame = 0
|
||||
|
||||
# Process each causal block
|
||||
for current_num_frames in block_sizes:
|
||||
# Extract current block from the full latents
|
||||
block_latents = current_noise_latents[:, start_index:start_index + current_num_frames, :, :, :]
|
||||
|
||||
timestep = torch.ones([batch_size, current_num_frames], device=noise.device, dtype=torch.int64) * current_timestep
|
||||
logger.info("timestep shape at initalization: %s", timestep.shape)
|
||||
|
||||
if not exit_flag:
|
||||
with torch.no_grad():
|
||||
# Build input kwargs
|
||||
# Process denoising timesteps for this block
|
||||
for step_idx in range(max_target_idx):
|
||||
current_timestep = self.denoising_step_list[step_idx]
|
||||
current_timestep_tensor = current_timestep * torch.ones(
|
||||
1, device=self.device, dtype=torch.long)
|
||||
|
||||
# Build input kwargs with KV cache support for this block
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
noisy_input, timestep, training_batch.conditional_dict, training_batch)
|
||||
block_latents, current_timestep_tensor,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
logger.info("timestep shape in generator: %s", timestep.shape)
|
||||
pred_flow = self.transformer(
|
||||
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
|
||||
timestep=training_batch_temp.input_kwargs['timestep'],
|
||||
encoder_hidden_states_image=training_batch_temp.input_kwargs.get('encoder_hidden_states_image'),
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=current_start_frame * self.frame_seq_length
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
denoised_pred = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep.flatten(),
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
|
||||
|
||||
next_timestep = self.denoising_step_list[index + 1]
|
||||
logger.info("denoised_pred shape before: %s", denoised_pred.shape)
|
||||
noisy_input = self.noise_scheduler.add_noise(
|
||||
denoised_pred.flatten(0, 1),
|
||||
torch.randn_like(denoised_pred.flatten(0, 1)),
|
||||
next_timestep * torch.ones([batch_size * current_num_frames], device=noise.device, dtype=torch.long)
|
||||
).unflatten(0, denoised_pred.shape[:2])
|
||||
logger.info("denoised_pred shape after: %s", denoised_pred.shape)
|
||||
else:
|
||||
# Final prediction with gradient control
|
||||
if current_start_frame < start_gradient_frame_index:
|
||||
logger.info("timestep shape in generator: %s", timestep.shape)
|
||||
with torch.no_grad():
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
noisy_input, timestep, training_batch.conditional_dict, training_batch)
|
||||
|
||||
# Add KV cache parameters for causal generation
|
||||
if hasattr(self.transformer, '_forward_inference'):
|
||||
# Use causal inference forward with KV cache
|
||||
pred_flow = self.transformer(
|
||||
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
|
||||
timestep=training_batch_temp.input_kwargs['timestep'],
|
||||
encoder_hidden_states_image=training_batch_temp.input_kwargs.get('encoder_hidden_states_image'),
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=current_start_frame * self.frame_seq_length
|
||||
kv_cache=kv_cache,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=current_start_frame * 1560, # TODO: remove hardcode
|
||||
cache_start=0
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
logger.info("timestep shape in generator: %s", timestep.shape)
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
noisy_input, timestep, training_batch.conditional_dict, training_batch)
|
||||
|
||||
pred_flow = self.transformer(
|
||||
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
|
||||
timestep=training_batch_temp.input_kwargs['timestep'],
|
||||
encoder_hidden_states_image=training_batch_temp.input_kwargs.get('encoder_hidden_states_image'),
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=current_start_frame * self.frame_seq_length
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
# Fallback to regular forward
|
||||
pred_flow = self.transformer(**training_batch_temp.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
pred_clean = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=block_latents.flatten(0, 1),
|
||||
timestep=current_timestep_tensor,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, pred_flow.shape[:2])
|
||||
|
||||
# Add noise for the next timestep
|
||||
if step_idx < max_target_idx - 1:
|
||||
next_timestep = self.denoising_step_list[step_idx + 1]
|
||||
next_timestep_tensor = next_timestep * torch.ones(
|
||||
1, device=self.device, dtype=torch.long)
|
||||
|
||||
# Generate noise for this block size
|
||||
block_noise_shape = (batch_size, current_num_frames, *self.video_latent_shape[2:])
|
||||
noise = torch.randn(block_noise_shape,
|
||||
device=self.device,
|
||||
dtype=pred_clean.dtype)
|
||||
if self.sp_world_size > 1:
|
||||
noise = rearrange(noise,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
|
||||
block_latents = self.noise_scheduler.add_noise(
|
||||
pred_clean.flatten(0, 1), noise.flatten(0, 1),
|
||||
next_timestep_tensor).unflatten(0, pred_clean.shape[:2])
|
||||
else:
|
||||
# Final step: use clean prediction
|
||||
block_latents = pred_clean
|
||||
|
||||
denoised_pred = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep.flatten(),
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
|
||||
break
|
||||
# Store the processed block result
|
||||
latent_copy = block_latents.clone()
|
||||
noise_latents.append(latent_copy)
|
||||
|
||||
# Update indices for next block
|
||||
current_start_frame += current_num_frames
|
||||
start_index += current_num_frames
|
||||
|
||||
# Reconstruct full latents from blocks
|
||||
if noise_latents:
|
||||
current_noise_latents = torch.cat(noise_latents, dim=1)
|
||||
|
||||
# Step 3.2: record the model's output
|
||||
output[:, current_start_frame:current_start_frame + current_num_frames] = denoised_pred
|
||||
|
||||
# Step 3.3: rerun with timestep zero to update the cache
|
||||
context_timestep = torch.ones_like(timestep) * self.context_noise
|
||||
logger.info("denoised_pred shape before: %s", denoised_pred.shape)
|
||||
denoised_pred = self.noise_scheduler.add_noise(
|
||||
denoised_pred.flatten(0, 1),
|
||||
torch.randn_like(denoised_pred.flatten(0, 1)),
|
||||
context_timestep * torch.ones([batch_size * current_num_frames], device=noise.device, dtype=torch.long)
|
||||
).unflatten(0, denoised_pred.shape[:2])
|
||||
logger.info("denoised_pred shape after: %s", denoised_pred.shape)
|
||||
|
||||
with torch.no_grad():
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
denoised_pred, context_timestep, training_batch.conditional_dict, training_batch)
|
||||
|
||||
self.transformer(
|
||||
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
|
||||
timestep=training_batch_temp.input_kwargs['timestep'],
|
||||
encoder_hidden_states_image=training_batch_temp.input_kwargs.get('encoder_hidden_states_image'),
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=current_start_frame * self.frame_seq_length
|
||||
)
|
||||
|
||||
# Step 3.4: update the start and end frame indices
|
||||
current_start_frame += current_num_frames
|
||||
|
||||
# Handle last 21 frames logic
|
||||
pred_image_or_video = output
|
||||
if num_input_frames > 0:
|
||||
pred_image_or_video = output[:, num_input_frames:]
|
||||
|
||||
# Slice last 21 frames if we generated more
|
||||
gradient_mask = None
|
||||
if pred_image_or_video.shape[1] > 21:
|
||||
with torch.no_grad():
|
||||
# Reencode to get image latent
|
||||
latent_to_decode = pred_image_or_video[:, :-20, ...]
|
||||
# Decode to video
|
||||
latent_to_decode = latent_to_decode.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
|
||||
|
||||
# Apply VAE scaling and shift factors
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latent_to_decode = latent_to_decode / self.vae.scaling_factor.to(latent_to_decode.device, latent_to_decode.dtype)
|
||||
else:
|
||||
latent_to_decode = latent_to_decode / self.vae.scaling_factor
|
||||
|
||||
if hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None:
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latent_to_decode += self.vae.shift_factor.to(latent_to_decode.device, latent_to_decode.dtype)
|
||||
else:
|
||||
latent_to_decode += self.vae.shift_factor
|
||||
|
||||
# Decode to pixels
|
||||
pixels = self.vae.decode(latent_to_decode)
|
||||
frame = pixels[:, :, -1:, :, :].to(dtype) # Last frame [B, C, 1, H, W]
|
||||
|
||||
# Encode frame back to get image latent
|
||||
image_latent = self.vae.encode(frame).mean.to(dtype)
|
||||
image_latent = image_latent.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
|
||||
|
||||
pred_image_or_video_last_21 = torch.cat([image_latent, pred_image_or_video[:, -20:, ...]], dim=1)
|
||||
# Step 6: Use the simulated noisy input for the final training step
|
||||
if noise_latent_index >= 0:
|
||||
assert noise_latent_index < len(
|
||||
self.denoising_step_list
|
||||
) - 1, "noise_latent_index is out of bounds"
|
||||
noisy_input = current_noise_latents
|
||||
else:
|
||||
pred_image_or_video_last_21 = pred_image_or_video
|
||||
|
||||
# Set up gradient mask if we generated more than minimum frames
|
||||
if num_generated_frames != min_num_frames:
|
||||
# Currently, we do not use gradient for the first chunk, since it contains image latents
|
||||
gradient_mask = torch.ones_like(pred_image_or_video_last_21, dtype=torch.bool)
|
||||
if self.independent_first_frame:
|
||||
gradient_mask[:, :1] = False
|
||||
else:
|
||||
gradient_mask[:, :self.num_frame_per_block] = False
|
||||
|
||||
# Apply gradient masking if needed
|
||||
final_output = pred_image_or_video_last_21.to(dtype)
|
||||
if gradient_mask is not None:
|
||||
# Apply gradient masking: detach frames that shouldn't contribute gradients
|
||||
final_output = torch.where(
|
||||
gradient_mask,
|
||||
pred_image_or_video_last_21, # Keep original values where gradient_mask is True
|
||||
pred_image_or_video_last_21.detach() # Detach where gradient_mask is False
|
||||
)
|
||||
noisy_input = current_noise_latents_copy
|
||||
|
||||
# Store visualization data
|
||||
training_batch.dmd_latent_vis_dict["generator_timestep"] = torch.tensor(
|
||||
self.denoising_step_list[exit_flags[0]], dtype=torch.float32, device=self.device)
|
||||
# Step 7: Final student prediction with selective gradient computation
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_input, target_timestep, training_batch.conditional_dict,
|
||||
training_batch)
|
||||
|
||||
# Apply gradient masking during final forward pass
|
||||
if gradient_mask is not None:
|
||||
# Create a custom forward function that applies gradient masking
|
||||
def masked_forward():
|
||||
pred_flow = self.transformer(**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=target_timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
|
||||
|
||||
# Apply gradient masking: detach frames that shouldn't contribute gradients
|
||||
masked_pred_video = pred_video.clone()
|
||||
masked_pred_video = torch.where(
|
||||
gradient_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1), # Broadcast to [B, F, 1, 1, 1]
|
||||
pred_video, # Keep original values where gradient_mask is True
|
||||
pred_video.detach() # Detach where gradient_mask is False
|
||||
)
|
||||
|
||||
return masked_pred_video
|
||||
|
||||
pred_video = masked_forward()
|
||||
else:
|
||||
pred_flow = self.transformer(**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=target_timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
|
||||
|
||||
training_batch.dmd_latent_vis_dict[
|
||||
"generator_timestep"] = target_timestep.float().detach()
|
||||
|
||||
# Store gradient mask information for debugging
|
||||
if gradient_mask is not None:
|
||||
training_batch.dmd_latent_vis_dict["gradient_mask"] = gradient_mask.float()
|
||||
training_batch.dmd_latent_vis_dict["num_generated_frames"] = torch.tensor(
|
||||
num_generated_frames, dtype=torch.float32, device=self.device)
|
||||
training_batch.dmd_latent_vis_dict["min_num_frames"] = torch.tensor(
|
||||
min_num_frames, dtype=torch.float32, device=self.device)
|
||||
start_gradient_frame_index = max(0, num_frames - gradient_mask_last_n_frames)
|
||||
training_batch.dmd_latent_vis_dict["start_gradient_frame_index"] = torch.tensor(
|
||||
start_gradient_frame_index, dtype=torch.float32, device=self.device)
|
||||
|
||||
self._reset_simulation_caches(self.kv_cache1, self.crossattn_cache)
|
||||
|
||||
return final_output if gradient_mask is not None else pred_image_or_video
|
||||
self._reset_simulation_caches(kv_cache, crossattn_cache)
|
||||
|
||||
return pred_video
|
||||
|
||||
def _initialize_simulation_caches(self, batch_size: int, dtype: torch.dtype, device: torch.device):
|
||||
"""Initialize KV cache and cross-attention cache for multi-step simulation."""
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import gc
|
||||
import dataclasses
|
||||
import math
|
||||
import os
|
||||
@@ -64,6 +65,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
train_dataloader: StatefulDataLoader
|
||||
train_loader_iter: Iterator[dict[str, Any]]
|
||||
current_epoch: int = 0
|
||||
train_transformer_2: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -99,6 +101,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
self.sp_world_size = self.sp_group.world_size
|
||||
self.local_rank = world_group.local_rank
|
||||
self.transformer = self.get_module("transformer")
|
||||
self.transformer_2 = self.get_module("transformer_2", None)
|
||||
self.seed = training_args.seed
|
||||
self.set_schemas()
|
||||
|
||||
@@ -111,10 +114,16 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
self.transformer,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
if self.transformer_2 is not None:
|
||||
self.transformer_2 = apply_activation_checkpointing(
|
||||
self.transformer_2,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
|
||||
noise_scheduler = self.modules["scheduler"]
|
||||
self.set_trainable()
|
||||
params_to_optimize = self.transformer.parameters()
|
||||
|
||||
params_to_optimize = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
# Parse betas from string format "beta1,beta2"
|
||||
@@ -142,6 +151,30 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
min_lr_ratio=training_args.min_lr_ratio,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
if self.transformer_2 is not None:
|
||||
# Ensure transformer_2 has trainable parameters before creating optimizer
|
||||
self.transformer_2.train()
|
||||
self.transformer_2.requires_grad_(True)
|
||||
params_to_optimize_2 = self.transformer_2.parameters()
|
||||
params_to_optimize_2 = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize_2))
|
||||
self.optimizer_2 = torch.optim.AdamW(
|
||||
params_to_optimize_2,
|
||||
lr=training_args.learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
self.lr_scheduler_2 = get_scheduler(
|
||||
training_args.lr_scheduler,
|
||||
optimizer=self.optimizer_2,
|
||||
num_warmup_steps=training_args.lr_warmup_steps,
|
||||
num_training_steps=training_args.max_train_steps,
|
||||
num_cycles=training_args.lr_num_cycles,
|
||||
power=training_args.lr_power,
|
||||
min_lr_ratio=training_args.min_lr_ratio,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
|
||||
training_args.data_path,
|
||||
@@ -156,7 +189,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
seed=self.seed)
|
||||
|
||||
self.noise_scheduler = noise_scheduler
|
||||
|
||||
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
|
||||
self.num_update_steps_per_epoch = math.ceil(
|
||||
len(self.train_dataloader) /
|
||||
training_args.gradient_accumulation_steps * training_args.sp_size /
|
||||
@@ -182,6 +215,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
def _prepare_training(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
self.transformer.train()
|
||||
self.optimizer.zero_grad()
|
||||
if self.transformer_2 is not None:
|
||||
self.transformer_2.train()
|
||||
self.optimizer_2.zero_grad()
|
||||
training_batch.total_loss = 0.0
|
||||
return training_batch
|
||||
|
||||
@@ -228,17 +264,8 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
generator=self.noise_gen_cuda,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype)
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=self.training_args.weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=self.noise_random_generator,
|
||||
logit_mean=self.training_args.logit_mean,
|
||||
logit_std=self.training_args.logit_std,
|
||||
mode_scale=self.training_args.mode_scale,
|
||||
)
|
||||
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
|
||||
timesteps = self.noise_scheduler.timesteps[indices].to(
|
||||
device=latents.device)
|
||||
timesteps = self._sample_timesteps(batch_size, latents.device)
|
||||
|
||||
if self.training_args.sp_size > 1:
|
||||
# Make sure that the timesteps are the same across all sp processes.
|
||||
sp_group = get_sp_group()
|
||||
@@ -261,6 +288,38 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
return training_batch
|
||||
|
||||
def _sample_timesteps(self, batch_size, device):
|
||||
# Determine which model to train based on the boundary timestep
|
||||
if (self.transformer_2 is not None and self.boundary_timestep is not None and
|
||||
torch.rand(1, generator=self.noise_random_generator).item() <= self.training_args.boundary_ratio):
|
||||
self.train_transformer_2 = True
|
||||
else:
|
||||
self.train_transformer_2 = False
|
||||
|
||||
# Broadcast the decision to all processes
|
||||
decision = torch.tensor(1.0 if self.train_transformer_2 else 0.0, device=self.device)
|
||||
dist.broadcast(decision, src=0)
|
||||
self.train_transformer_2 = decision.item() == 1.0
|
||||
|
||||
# Sample u from the appropriate range
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=self.training_args.weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=self.noise_random_generator,
|
||||
logit_mean=self.training_args.logit_mean,
|
||||
logit_std=self.training_args.logit_std,
|
||||
mode_scale=self.training_args.mode_scale,
|
||||
)
|
||||
|
||||
boundary_ratio = self.training_args.boundary_ratio
|
||||
if self.train_transformer_2:
|
||||
u = (1 - boundary_ratio) + u * boundary_ratio # min: 1 - boundary_ratio, max: 1
|
||||
else:
|
||||
u = u * (1 - boundary_ratio) # min: 0, max: 1 - boundary_ratio
|
||||
|
||||
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
|
||||
return self.noise_scheduler.timesteps[indices].to(device=device)
|
||||
|
||||
def _build_attention_metadata(
|
||||
self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
latents_shape = training_batch.raw_latent_shape
|
||||
@@ -311,11 +370,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# [1000.0],
|
||||
# device=training_batch.noisy_model_input.device,
|
||||
# dtype=torch.bfloat16)
|
||||
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.current_timestep,
|
||||
attn_metadata=training_batch.attn_metadata):
|
||||
model_pred = self.transformer(**input_kwargs)
|
||||
model_pred = current_model(**input_kwargs)
|
||||
if self.training_args.precondition_outputs:
|
||||
assert training_batch.sigmas is not None
|
||||
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
|
||||
@@ -346,7 +406,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# the following:
|
||||
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
if max_grad_norm is not None:
|
||||
model_parts = [self.transformer]
|
||||
# Only clip gradients for the model that is currently training
|
||||
if self.train_transformer_2 and self.transformer_2 is not None:
|
||||
model_parts = [self.transformer_2]
|
||||
else:
|
||||
model_parts = [self.transformer]
|
||||
|
||||
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
|
||||
[p for m in model_parts for p in m.parameters()],
|
||||
max_grad_norm,
|
||||
@@ -391,9 +456,14 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
training_batch = self._clip_grad_norm(training_batch)
|
||||
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
# Only step the optimizer and scheduler for the model that is currently training
|
||||
if self.train_transformer_2 and self.transformer_2 is not None:
|
||||
self.optimizer_2.step()
|
||||
self.lr_scheduler_2.step()
|
||||
else:
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
training_batch.total_loss = training_batch.total_loss
|
||||
training_batch.grad_norm = training_batch.grad_norm
|
||||
return training_batch
|
||||
@@ -425,6 +495,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
logger.info("Starting training with %s B trainable parameters",
|
||||
round(num_trainable_params / 1e9, 3))
|
||||
|
||||
if getattr(self, "transformer_2", None) is not None:
|
||||
num_trainable_params = _get_trainable_params(self.transformer_2)
|
||||
logger.info("Transformer 2: Starting training with %s B trainable parameters",
|
||||
round(num_trainable_params / 1e9, 3))
|
||||
|
||||
# Set random seeds for deterministic training
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
@@ -445,7 +520,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
self._log_training_info()
|
||||
|
||||
self._log_validation(self.transformer, self.training_args,
|
||||
self._log_validation(self.training_args,
|
||||
self.init_steps)
|
||||
|
||||
# Train!
|
||||
@@ -507,7 +582,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
self.transformer.train()
|
||||
self.sp_group.barrier()
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
self._log_validation(self.transformer, self.training_args, step)
|
||||
self._log_validation(self.training_args, step)
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
trainable_params = round(
|
||||
_get_trainable_params(self.transformer) / 1e9, 3)
|
||||
@@ -588,12 +663,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
return batch
|
||||
|
||||
@torch.no_grad()
|
||||
def _log_validation(self, transformer, training_args, global_step) -> None:
|
||||
def _log_validation(self, training_args, global_step) -> None:
|
||||
"""
|
||||
Generate a validation video and log it to wandb to check the quality during training.
|
||||
"""
|
||||
training_args.inference_mode = True
|
||||
training_args.dit_cpu_offload = True
|
||||
training_args.dit_cpu_offload = False
|
||||
if not training_args.log_validation:
|
||||
return
|
||||
if self.validation_pipeline is None:
|
||||
@@ -615,7 +690,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
batch_size=None,
|
||||
num_workers=0)
|
||||
|
||||
transformer.eval()
|
||||
self.transformer.eval()
|
||||
if getattr(self, "transformer_2", None) is not None:
|
||||
self.transformer_2.eval()
|
||||
|
||||
validation_steps = training_args.validation_sampling_steps.split(",")
|
||||
validation_steps = [int(step) for step in validation_steps]
|
||||
@@ -707,4 +784,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
# Re-enable gradients for training
|
||||
training_args.inference_mode = False
|
||||
transformer.train()
|
||||
self.transformer.train()
|
||||
if getattr(self, "transformer_2", None) is not None:
|
||||
self.transformer_2.train()
|
||||
|
||||
@@ -0,0 +1,797 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import dataclasses
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import deque
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torchvision
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from einops import rearrange
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionMetadataBuilder)
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset import build_parquet_map_style_dataloader
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
from fastvideo.dataset.validation_dataset import ValidationDataset
|
||||
from fastvideo.distributed import (cleanup_dist_env_and_memory,
|
||||
get_local_torch_device, get_sp_group,
|
||||
get_world_group)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines import (ComposedPipelineBase, ForwardBatch,
|
||||
LoRAPipeline, TrainingBatch)
|
||||
from fastvideo.training.activation_checkpoint import (
|
||||
apply_activation_checkpointing)
|
||||
from fastvideo.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
compute_density_for_timestep_sampling, get_scheduler, get_sigmas,
|
||||
load_checkpoint, normalize_dit_input, save_checkpoint,
|
||||
shard_latents_across_sp)
|
||||
from fastvideo.utils import is_vsa_available, set_random_seed, shallow_asdict
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _get_trainable_params(model: torch.nn.Module) -> int:
|
||||
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
|
||||
|
||||
class TrainingPipeline(LoRAPipeline, ABC):
|
||||
"""
|
||||
A pipeline for training a model. All training pipelines should inherit from this class.
|
||||
All reusable components and code should be implemented in this class.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer"]
|
||||
validation_pipeline: ComposedPipelineBase
|
||||
train_dataloader: StatefulDataLoader
|
||||
train_loader_iter: Iterator[dict[str, Any]]
|
||||
current_epoch: int = 0
|
||||
train_transformer_2: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_path: str,
|
||||
fastvideo_args: TrainingArgs,
|
||||
required_config_modules: list[str] | None = None,
|
||||
loaded_modules: dict[str, torch.nn.Module] | None = None) -> None:
|
||||
fastvideo_args.inference_mode = False
|
||||
self.lora_training = fastvideo_args.lora_training
|
||||
if self.lora_training and fastvideo_args.lora_rank is None:
|
||||
raise ValueError("lora rank must be set when using lora training")
|
||||
|
||||
set_random_seed(fastvideo_args.seed) # for lora param init
|
||||
super().__init__(model_path, fastvideo_args, required_config_modules,
|
||||
loaded_modules) # type: ignore
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
raise RuntimeError(
|
||||
"create_pipeline_stages should not be called for training pipeline")
|
||||
|
||||
def set_schemas(self) -> None:
|
||||
self.train_dataset_schema = pyarrow_schema_t2v
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing training pipeline...")
|
||||
self.device = get_local_torch_device()
|
||||
self.training_args = training_args
|
||||
world_group = get_world_group()
|
||||
self.world_size = world_group.world_size
|
||||
self.global_rank = world_group.rank
|
||||
self.sp_group = get_sp_group()
|
||||
self.rank_in_sp_group = self.sp_group.rank_in_group
|
||||
self.sp_world_size = self.sp_group.world_size
|
||||
self.local_rank = world_group.local_rank
|
||||
self.transformer = self.get_module("transformer")
|
||||
self.transformer_2 = self.get_module("transformer_2", None)
|
||||
self.seed = training_args.seed
|
||||
self.set_schemas()
|
||||
|
||||
# Set random seeds for deterministic training
|
||||
assert self.seed is not None, "seed must be set"
|
||||
set_random_seed(self.seed)
|
||||
self.transformer.train()
|
||||
if training_args.enable_gradient_checkpointing_type is not None:
|
||||
self.transformer = apply_activation_checkpointing(
|
||||
self.transformer,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
if self.transformer_2 is not None:
|
||||
self.transformer_2 = apply_activation_checkpointing(
|
||||
self.transformer_2,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
|
||||
noise_scheduler = self.modules["scheduler"]
|
||||
self.set_trainable()
|
||||
params_to_optimize = self.transformer.parameters()
|
||||
|
||||
params_to_optimize = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
self.optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
lr=training_args.learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
self.init_steps = 0
|
||||
logger.info("optimizer: %s", self.optimizer)
|
||||
|
||||
self.lr_scheduler = get_scheduler(
|
||||
training_args.lr_scheduler,
|
||||
optimizer=self.optimizer,
|
||||
num_warmup_steps=training_args.lr_warmup_steps,
|
||||
num_training_steps=training_args.max_train_steps,
|
||||
num_cycles=training_args.lr_num_cycles,
|
||||
power=training_args.lr_power,
|
||||
min_lr_ratio=training_args.min_lr_ratio,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
if self.transformer_2 is not None:
|
||||
# Ensure transformer_2 has trainable parameters before creating optimizer
|
||||
self.transformer_2.train()
|
||||
self.transformer_2.requires_grad_(True)
|
||||
params_to_optimize_2 = self.transformer_2.parameters()
|
||||
params_to_optimize_2 = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize_2))
|
||||
self.optimizer_2 = torch.optim.AdamW(
|
||||
params_to_optimize_2,
|
||||
lr=training_args.learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
self.lr_scheduler_2 = get_scheduler(
|
||||
training_args.lr_scheduler,
|
||||
optimizer=self.optimizer_2,
|
||||
num_warmup_steps=training_args.lr_warmup_steps,
|
||||
num_training_steps=training_args.max_train_steps,
|
||||
num_cycles=training_args.lr_num_cycles,
|
||||
power=training_args.lr_power,
|
||||
min_lr_ratio=training_args.min_lr_ratio,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
|
||||
training_args.data_path,
|
||||
training_args.train_batch_size,
|
||||
parquet_schema=self.train_dataset_schema,
|
||||
num_data_workers=training_args.dataloader_num_workers,
|
||||
cfg_rate=training_args.training_cfg_rate,
|
||||
drop_last=True,
|
||||
text_padding_length=training_args.pipeline_config.
|
||||
text_encoder_configs[0].arch_config.
|
||||
text_len, # type: ignore[attr-defined]
|
||||
seed=self.seed)
|
||||
|
||||
self.noise_scheduler = noise_scheduler
|
||||
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
|
||||
self.num_update_steps_per_epoch = math.ceil(
|
||||
len(self.train_dataloader) /
|
||||
training_args.gradient_accumulation_steps * training_args.sp_size /
|
||||
training_args.train_sp_batch_size)
|
||||
self.num_train_epochs = math.ceil(training_args.max_train_steps /
|
||||
self.num_update_steps_per_epoch)
|
||||
|
||||
# TODO(will): is there a cleaner way to track epochs?
|
||||
self.current_epoch = 0
|
||||
|
||||
if self.global_rank == 0:
|
||||
project = training_args.tracker_project_name or "fastvideo"
|
||||
wandb_config = dataclasses.asdict(training_args)
|
||||
wandb.init(project=project,
|
||||
config=wandb_config,
|
||||
name=training_args.wandb_run_name)
|
||||
|
||||
@abstractmethod
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
raise NotImplementedError(
|
||||
"Training pipelines must implement this method")
|
||||
|
||||
def _prepare_training(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
# At the beginning, disable training for all models
|
||||
self._disable_training(self.transformer, self.optimizer)
|
||||
if self.transformer_2 is not None:
|
||||
self._disable_training(self.transformer_2, self.optimizer_2)
|
||||
|
||||
training_batch.total_loss = 0.0
|
||||
return training_batch
|
||||
|
||||
def _enable_training(self, model: torch.nn.Module, optimizer: torch.optim.Optimizer) -> None:
|
||||
"""Enable training mode and gradients for the specified model."""
|
||||
for param in model.parameters():
|
||||
param.requires_grad = True
|
||||
model.train()
|
||||
optimizer.zero_grad()
|
||||
|
||||
def _disable_training(self, model: torch.nn.Module, optimizer: torch.optim.Optimizer) -> None:
|
||||
"""Disable training mode and gradients for the specified model."""
|
||||
for param in model.parameters():
|
||||
param.requires_grad = False
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
batch = next(self.train_loader_iter, None) # type: ignore
|
||||
if batch is None:
|
||||
self.current_epoch += 1
|
||||
logger.info("Starting epoch %s", self.current_epoch)
|
||||
# Reset iterator for next epoch
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
# Get first batch of new epoch
|
||||
batch = next(self.train_loader_iter)
|
||||
|
||||
# latents, encoder_hidden_states, encoder_attention_mask, infos = batch
|
||||
latents = batch['vae_latent']
|
||||
latents = latents[:, :, :self.training_args.num_latent_t]
|
||||
encoder_hidden_states = batch['text_embedding']
|
||||
encoder_attention_mask = batch['text_attention_mask']
|
||||
infos = batch['info_list']
|
||||
|
||||
training_batch.latents = latents.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
training_batch.encoder_hidden_states = encoder_hidden_states.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch
|
||||
|
||||
def _normalize_dit_input(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
# TODO(will): support other models
|
||||
training_batch.latents = normalize_dit_input('wan',
|
||||
training_batch.latents,
|
||||
self.get_module("vae"))
|
||||
return training_batch
|
||||
|
||||
def _prepare_dit_inputs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
latents = training_batch.latents
|
||||
batch_size = latents.shape[0]
|
||||
noise = torch.randn(latents.shape,
|
||||
generator=self.noise_gen_cuda,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype)
|
||||
timesteps = self._sample_timesteps(batch_size, latents.device)
|
||||
|
||||
# Enable training for the model that will be trained next and disable the other
|
||||
if self.train_transformer_2:
|
||||
self._enable_training(self.transformer_2, self.optimizer_2)
|
||||
self._disable_training(self.transformer, self.optimizer)
|
||||
else:
|
||||
self._enable_training(self.transformer, self.optimizer)
|
||||
if self.transformer_2 is not None:
|
||||
self._disable_training(self.transformer_2, self.optimizer_2)
|
||||
|
||||
if self.training_args.sp_size > 1:
|
||||
# Make sure that the timesteps are the same across all sp processes.
|
||||
sp_group = get_sp_group()
|
||||
sp_group.broadcast(timesteps, src=0)
|
||||
sigmas = get_sigmas(
|
||||
self.noise_scheduler,
|
||||
latents.device,
|
||||
timesteps,
|
||||
n_dim=latents.ndim,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
noisy_model_input = (1.0 -
|
||||
sigmas) * training_batch.latents + sigmas * noise
|
||||
|
||||
training_batch.noisy_model_input = noisy_model_input
|
||||
training_batch.timesteps = timesteps
|
||||
training_batch.sigmas = sigmas
|
||||
training_batch.noise = noise
|
||||
training_batch.raw_latent_shape = training_batch.latents.shape
|
||||
|
||||
return training_batch
|
||||
|
||||
def _sample_timesteps(self, batch_size, device):
|
||||
# Determine which model to train based on the boundary timestep
|
||||
if (self.transformer_2 is not None and self.boundary_timestep is not None and
|
||||
torch.rand(1, generator=self.noise_random_generator).item() > self.training_args.boundary_ratio):
|
||||
self.train_transformer_2 = True
|
||||
else:
|
||||
self.train_transformer_2 = False
|
||||
|
||||
# Broadcast the decision to all processes
|
||||
decision = torch.tensor(1.0 if self.train_transformer_2 else 0.0, device=self.device)
|
||||
dist.broadcast(decision, src=0)
|
||||
self.train_transformer_2 = decision.item() == 1.0
|
||||
|
||||
# Sample u from the appropriate range
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=self.training_args.weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=self.noise_random_generator,
|
||||
logit_mean=self.training_args.logit_mean,
|
||||
logit_std=self.training_args.logit_std,
|
||||
mode_scale=self.training_args.mode_scale,
|
||||
)
|
||||
|
||||
boundary_ratio = self.training_args.boundary_ratio
|
||||
if self.train_transformer_2:
|
||||
u = boundary_ratio + u * (1.0 - boundary_ratio)
|
||||
else:
|
||||
u = u * boundary_ratio
|
||||
|
||||
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
|
||||
return self.noise_scheduler.timesteps[indices].to(device=device)
|
||||
|
||||
def _build_attention_metadata(
|
||||
self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
latents_shape = training_batch.raw_latent_shape
|
||||
patch_size = self.training_args.pipeline_config.dit_config.patch_size
|
||||
current_vsa_sparsity = training_batch.current_vsa_sparsity
|
||||
assert latents_shape is not None
|
||||
assert training_batch.timesteps is not None
|
||||
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
|
||||
training_batch.attn_metadata = VideoSparseAttentionMetadataBuilder( # type: ignore
|
||||
).build( # type: ignore
|
||||
raw_latent_shape=latents_shape[2:5],
|
||||
current_timestep=training_batch.timesteps,
|
||||
patch_size=patch_size,
|
||||
VSA_sparsity=current_vsa_sparsity,
|
||||
device=get_local_torch_device())
|
||||
else:
|
||||
training_batch.attn_metadata = None
|
||||
|
||||
return training_batch
|
||||
|
||||
def _build_input_kwargs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
training_batch.input_kwargs = {
|
||||
"hidden_states":
|
||||
training_batch.noisy_model_input,
|
||||
"encoder_hidden_states":
|
||||
training_batch.encoder_hidden_states,
|
||||
"timestep":
|
||||
training_batch.timesteps.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16),
|
||||
"encoder_attention_mask":
|
||||
training_batch.encoder_attention_mask,
|
||||
"return_dict":
|
||||
False,
|
||||
}
|
||||
return training_batch
|
||||
|
||||
def _transformer_forward_and_compute_loss(
|
||||
self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
|
||||
assert training_batch.attn_metadata is not None
|
||||
else:
|
||||
assert training_batch.attn_metadata is None
|
||||
input_kwargs = training_batch.input_kwargs
|
||||
|
||||
# if 'hunyuan' in self.training_args.model_type:
|
||||
# input_kwargs["guidance"] = torch.tensor(
|
||||
# [1000.0],
|
||||
# device=training_batch.noisy_model_input.device,
|
||||
# dtype=torch.bfloat16)
|
||||
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.current_timestep,
|
||||
attn_metadata=training_batch.attn_metadata):
|
||||
model_pred = current_model(**input_kwargs)
|
||||
if self.training_args.precondition_outputs:
|
||||
assert training_batch.sigmas is not None
|
||||
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
|
||||
assert training_batch.latents is not None
|
||||
assert training_batch.noise is not None
|
||||
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
|
||||
|
||||
# make sure no implicit broadcasting happens
|
||||
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
|
||||
loss = (torch.mean((model_pred.float() - target.float())**2) /
|
||||
self.training_args.gradient_accumulation_steps)
|
||||
|
||||
loss.backward()
|
||||
avg_loss = loss.detach().clone()
|
||||
|
||||
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
|
||||
# local_main_process_only=False)
|
||||
world_group = get_world_group()
|
||||
world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
|
||||
training_batch.total_loss += avg_loss.item()
|
||||
|
||||
return training_batch
|
||||
|
||||
def _clip_grad_norm(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
max_grad_norm = self.training_args.max_grad_norm
|
||||
|
||||
# TODO(will): perhaps move this into transformer api so that we can do
|
||||
# the following:
|
||||
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
if max_grad_norm is not None:
|
||||
# Only clip gradients for the model that is currently training
|
||||
if self.train_transformer_2 and self.transformer_2 is not None:
|
||||
model_parts = [self.transformer_2]
|
||||
else:
|
||||
model_parts = [self.transformer]
|
||||
|
||||
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
|
||||
[p for m in model_parts for p in m.parameters()],
|
||||
max_grad_norm,
|
||||
foreach=None,
|
||||
)
|
||||
assert grad_norm is not float('nan') or grad_norm is not float(
|
||||
'inf')
|
||||
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
|
||||
else:
|
||||
grad_norm = 0.0
|
||||
training_batch.grad_norm = grad_norm
|
||||
return training_batch
|
||||
|
||||
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
training_batch = self._prepare_training(training_batch)
|
||||
|
||||
for _ in range(self.training_args.gradient_accumulation_steps):
|
||||
training_batch = self._get_next_batch(training_batch)
|
||||
|
||||
# Normalize DIT input
|
||||
training_batch = self._normalize_dit_input(training_batch)
|
||||
# Create noisy model input
|
||||
training_batch = self._prepare_dit_inputs(training_batch)
|
||||
|
||||
# Shard latents across sp groups
|
||||
training_batch.latents = shard_latents_across_sp(
|
||||
training_batch.latents,
|
||||
num_latent_t=self.training_args.num_latent_t)
|
||||
# shard noisy_model_input to match
|
||||
training_batch.noisy_model_input = shard_latents_across_sp(
|
||||
training_batch.noisy_model_input,
|
||||
num_latent_t=self.training_args.num_latent_t)
|
||||
# shard noise to match latents
|
||||
training_batch.noise = shard_latents_across_sp(
|
||||
training_batch.noise,
|
||||
num_latent_t=self.training_args.num_latent_t)
|
||||
|
||||
training_batch = self._build_attention_metadata(training_batch)
|
||||
training_batch = self._build_input_kwargs(training_batch)
|
||||
training_batch = self._transformer_forward_and_compute_loss(
|
||||
training_batch)
|
||||
|
||||
training_batch = self._clip_grad_norm(training_batch)
|
||||
|
||||
# Only step the optimizer and scheduler for the model that is currently training
|
||||
if self.train_transformer_2 and self.transformer_2 is not None:
|
||||
self.optimizer_2.step()
|
||||
self.lr_scheduler_2.step()
|
||||
else:
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
training_batch.total_loss = training_batch.total_loss
|
||||
training_batch.grad_norm = training_batch.grad_norm
|
||||
return training_batch
|
||||
|
||||
def _resume_from_checkpoint(self) -> None:
|
||||
logger.info("Loading checkpoint from %s",
|
||||
self.training_args.resume_from_checkpoint)
|
||||
resumed_step = load_checkpoint(
|
||||
self.transformer, self.global_rank,
|
||||
self.training_args.resume_from_checkpoint, self.optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
self.noise_random_generator)
|
||||
if resumed_step > 0:
|
||||
self.init_steps = resumed_step
|
||||
logger.info("Successfully resumed from step %s", resumed_step)
|
||||
else:
|
||||
logger.warning("Failed to load checkpoint, starting from step 0")
|
||||
self.init_steps = 0
|
||||
|
||||
def train(self) -> None:
|
||||
assert self.seed is not None, "seed must be set"
|
||||
set_random_seed(self.seed + self.global_rank)
|
||||
logger.info('rank: %s: start training',
|
||||
self.global_rank,
|
||||
local_main_process_only=False)
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
num_trainable_params = _get_trainable_params(self.transformer)
|
||||
logger.info("Starting training with %s B trainable parameters",
|
||||
round(num_trainable_params / 1e9, 3))
|
||||
|
||||
# Set random seeds for deterministic training
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
self.noise_gen_cuda = torch.Generator(device="cuda").manual_seed(
|
||||
self.seed)
|
||||
self.validation_random_generator = torch.Generator(
|
||||
device="cpu").manual_seed(self.seed)
|
||||
logger.info("Initialized random seeds with seed: %s", self.seed)
|
||||
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
|
||||
if self.training_args.resume_from_checkpoint:
|
||||
self._resume_from_checkpoint()
|
||||
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
|
||||
step_times: deque[float] = deque(maxlen=100)
|
||||
|
||||
self._log_training_info()
|
||||
|
||||
self._log_validation(self.transformer, self.training_args,
|
||||
self.init_steps)
|
||||
|
||||
# Train!
|
||||
progress_bar = tqdm(
|
||||
range(0, self.training_args.max_train_steps),
|
||||
initial=self.init_steps,
|
||||
desc="Steps",
|
||||
# Only show the progress bar once on each machine.
|
||||
disable=self.local_rank > 0,
|
||||
)
|
||||
for step in range(self.init_steps + 1,
|
||||
self.training_args.max_train_steps + 1):
|
||||
start_time = time.perf_counter()
|
||||
if vsa_available:
|
||||
vsa_sparsity = self.training_args.VSA_sparsity
|
||||
vsa_decay_rate = self.training_args.VSA_decay_rate
|
||||
vsa_decay_interval_steps = self.training_args.VSA_decay_interval_steps
|
||||
current_decay_times = min(step // vsa_decay_interval_steps,
|
||||
vsa_sparsity // vsa_decay_rate)
|
||||
current_vsa_sparsity = current_decay_times * vsa_decay_rate
|
||||
else:
|
||||
current_vsa_sparsity = 0.0
|
||||
|
||||
training_batch = TrainingBatch()
|
||||
training_batch.current_timestep = step
|
||||
training_batch.current_vsa_sparsity = current_vsa_sparsity
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
loss = training_batch.total_loss
|
||||
grad_norm = training_batch.grad_norm
|
||||
|
||||
step_time = time.perf_counter() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"loss": f"{loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if self.global_rank == 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss": loss,
|
||||
"learning_rate": self.lr_scheduler.get_last_lr()[0],
|
||||
"step_time": step_time,
|
||||
"avg_step_time": avg_step_time,
|
||||
"grad_norm": grad_norm,
|
||||
"vsa_sparsity": current_vsa_sparsity,
|
||||
},
|
||||
step=step,
|
||||
)
|
||||
if step % self.training_args.checkpointing_steps == 0:
|
||||
save_checkpoint(self.transformer, self.global_rank,
|
||||
self.training_args.output_dir, step,
|
||||
self.optimizer, self.train_dataloader,
|
||||
self.lr_scheduler, self.noise_random_generator)
|
||||
self.transformer.train()
|
||||
self.sp_group.barrier()
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
self._log_validation(self.transformer, self.training_args, step)
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
trainable_params = round(
|
||||
_get_trainable_params(self.transformer) / 1e9, 3)
|
||||
logger.info(
|
||||
"GPU memory usage after validation: %s MB, trainable params: %sB",
|
||||
gpu_memory_usage, trainable_params)
|
||||
|
||||
wandb.finish()
|
||||
save_checkpoint(self.transformer, self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
self.training_args.max_train_steps, self.optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
self.noise_random_generator)
|
||||
|
||||
if get_sp_group():
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
def _log_training_info(self) -> None:
|
||||
total_batch_size = (self.world_size *
|
||||
self.training_args.gradient_accumulation_steps /
|
||||
self.training_args.sp_size *
|
||||
self.training_args.train_sp_batch_size)
|
||||
logger.info("***** Running training *****")
|
||||
logger.info(" Num examples = %s", len(self.train_dataset))
|
||||
logger.info(" Dataloader size = %s", len(self.train_dataloader))
|
||||
logger.info(" Num Epochs = %s", self.num_train_epochs)
|
||||
logger.info(" Resume training from step %s",
|
||||
self.init_steps) # type: ignore
|
||||
logger.info(" Instantaneous batch size per device = %s",
|
||||
self.training_args.train_batch_size)
|
||||
logger.info(
|
||||
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
|
||||
total_batch_size)
|
||||
logger.info(" Gradient Accumulation steps = %s",
|
||||
self.training_args.gradient_accumulation_steps)
|
||||
logger.info(" Total optimization steps = %s",
|
||||
self.training_args.max_train_steps)
|
||||
logger.info(" Total training parameters per FSDP shard = %s B",
|
||||
round(_get_trainable_params(self.transformer) / 1e9, 3))
|
||||
# print dtype
|
||||
logger.info(" Master weight dtype: %s",
|
||||
self.transformer.parameters().__next__().dtype)
|
||||
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
logger.info("GPU memory usage before train_one_step: %s MB",
|
||||
gpu_memory_usage)
|
||||
logger.info("VSA validation sparsity: %s",
|
||||
self.training_args.VSA_sparsity)
|
||||
|
||||
def _prepare_validation_batch(self, sampling_param: SamplingParam,
|
||||
training_args: TrainingArgs,
|
||||
validation_batch: dict[str, Any],
|
||||
num_inference_steps: int) -> ForwardBatch:
|
||||
sampling_param.prompt = validation_batch['prompt']
|
||||
sampling_param.height = training_args.num_height
|
||||
sampling_param.width = training_args.num_width
|
||||
sampling_param.num_inference_steps = num_inference_steps
|
||||
sampling_param.data_type = "video"
|
||||
assert self.seed is not None
|
||||
sampling_param.seed = self.seed
|
||||
|
||||
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
|
||||
sampling_param.height // 8, sampling_param.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = (training_args.num_latent_t -
|
||||
1) * temporal_compression_factor + 1
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=self.validation_random_generator,
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
VSA_sparsity=training_args.VSA_sparsity,
|
||||
)
|
||||
|
||||
return batch
|
||||
|
||||
@torch.no_grad()
|
||||
def _log_validation(self, transformer, training_args, global_step) -> None:
|
||||
"""
|
||||
Generate a validation video and log it to wandb to check the quality during training.
|
||||
"""
|
||||
training_args.inference_mode = True
|
||||
training_args.dit_cpu_offload = True
|
||||
if not training_args.log_validation:
|
||||
return
|
||||
if self.validation_pipeline is None:
|
||||
raise ValueError("Validation pipeline is not set")
|
||||
|
||||
logger.info("Starting validation")
|
||||
|
||||
# Create sampling parameters if not provided
|
||||
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
|
||||
|
||||
# Prepare validation prompts
|
||||
logger.info('rank: %s: fastvideo_args.validation_dataset_file: %s',
|
||||
self.global_rank,
|
||||
training_args.validation_dataset_file,
|
||||
local_main_process_only=False)
|
||||
validation_dataset = ValidationDataset(
|
||||
training_args.validation_dataset_file)
|
||||
validation_dataloader = DataLoader(validation_dataset,
|
||||
batch_size=None,
|
||||
num_workers=0)
|
||||
|
||||
transformer.eval()
|
||||
|
||||
validation_steps = training_args.validation_sampling_steps.split(",")
|
||||
validation_steps = [int(step) for step in validation_steps]
|
||||
validation_steps = [step for step in validation_steps if step > 0]
|
||||
# Log validation results for this step
|
||||
world_group = get_world_group()
|
||||
num_sp_groups = world_group.world_size // self.sp_group.world_size
|
||||
|
||||
# Process each validation prompt for each validation step
|
||||
for num_inference_steps in validation_steps:
|
||||
logger.info("rank: %s: num_inference_steps: %s",
|
||||
self.global_rank,
|
||||
num_inference_steps,
|
||||
local_main_process_only=False)
|
||||
step_videos: list[np.ndarray] = []
|
||||
step_captions: list[str] = []
|
||||
|
||||
for validation_batch in validation_dataloader:
|
||||
batch = self._prepare_validation_batch(sampling_param,
|
||||
training_args,
|
||||
validation_batch,
|
||||
num_inference_steps)
|
||||
logger.info("rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
|
||||
self.global_rank,
|
||||
self.rank_in_sp_group,
|
||||
batch.prompt,
|
||||
local_main_process_only=False)
|
||||
|
||||
assert batch.prompt is not None and isinstance(
|
||||
batch.prompt, str)
|
||||
step_captions.append(batch.prompt)
|
||||
|
||||
# Run validation inference
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, training_args)
|
||||
samples = output_batch.output
|
||||
|
||||
if self.rank_in_sp_group != 0:
|
||||
continue
|
||||
|
||||
# Process outputs
|
||||
video = rearrange(samples, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in video:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
step_videos.append(frames)
|
||||
|
||||
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
|
||||
# results to global rank 0
|
||||
if self.rank_in_sp_group == 0:
|
||||
if self.global_rank == 0:
|
||||
# Global rank 0 collects results from all sp_group leaders
|
||||
all_videos = step_videos # Start with own results
|
||||
all_captions = step_captions
|
||||
|
||||
# Receive from other sp_group leaders
|
||||
for sp_group_idx in range(1, num_sp_groups):
|
||||
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
|
||||
recv_videos = world_group.recv_object(src=src_rank)
|
||||
recv_captions = world_group.recv_object(src=src_rank)
|
||||
all_videos.extend(recv_videos)
|
||||
all_captions.extend(recv_captions)
|
||||
|
||||
video_filenames = []
|
||||
for i, (video, caption) in enumerate(
|
||||
zip(all_videos, all_captions, strict=True)):
|
||||
os.makedirs(training_args.output_dir, exist_ok=True)
|
||||
filename = os.path.join(
|
||||
training_args.output_dir,
|
||||
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
|
||||
)
|
||||
imageio.mimsave(filename, video, fps=sampling_param.fps)
|
||||
video_filenames.append(filename)
|
||||
|
||||
logs = {
|
||||
f"validation_videos_{num_inference_steps}_steps": [
|
||||
wandb.Video(filename, caption=caption)
|
||||
for filename, caption in zip(
|
||||
video_filenames, all_captions, strict=True)
|
||||
]
|
||||
}
|
||||
wandb.log(logs, step=global_step)
|
||||
else:
|
||||
# Other sp_group leaders send their results to global rank 0
|
||||
world_group.send_object(step_videos, dst=0)
|
||||
world_group.send_object(step_captions, dst=0)
|
||||
|
||||
# Re-enable gradients for training
|
||||
training_args.inference_mode = False
|
||||
transformer.train()
|
||||
@@ -34,7 +34,7 @@ class WanTrainingPipeline(TrainingPipeline):
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
assert training_args.dit_cpu_offload == False
|
||||
args_copy.inference_mode = True
|
||||
validation_pipeline = WanPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
@@ -42,12 +42,13 @@ class WanTrainingPipeline(TrainingPipeline):
|
||||
inference_mode=True,
|
||||
loaded_modules={
|
||||
"transformer": self.get_module("transformer"),
|
||||
"transformer_2": self.get_module("transformer_2")
|
||||
},
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
pin_cpu_memory=training_args.pin_cpu_memory,
|
||||
dit_cpu_offload=True)
|
||||
dit_cpu_offload=training_args.dit_cpu_offload)
|
||||
|
||||
self.validation_pipeline = validation_pipeline
|
||||
|
||||
|
||||
@@ -812,3 +812,66 @@ def set_random_seed(seed: int) -> None:
|
||||
@lru_cache(maxsize=1)
|
||||
def is_vsa_available() -> bool:
|
||||
return importlib.util.find_spec("vsa") is not None
|
||||
|
||||
|
||||
# adapted from: https://github.com/Wan-Video/Wan2.2/blob/main/wan/utils/utils.py
|
||||
def masks_like(tensor,
|
||||
zero=False,
|
||||
generator=None,
|
||||
p=0.2) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
|
||||
assert isinstance(tensor, list)
|
||||
out1 = [torch.ones(u.shape, dtype=u.dtype, device=u.device) for u in tensor]
|
||||
|
||||
out2 = [torch.ones(u.shape, dtype=u.dtype, device=u.device) for u in tensor]
|
||||
|
||||
if zero:
|
||||
if generator is not None:
|
||||
for u, v in zip(out1, out2, strict=False):
|
||||
random_num = torch.rand(1,
|
||||
generator=generator,
|
||||
device=generator.device).item()
|
||||
if random_num < p:
|
||||
u[:, 0] = torch.normal(mean=-3.5,
|
||||
std=0.5,
|
||||
size=(1, ),
|
||||
device=u.device,
|
||||
generator=generator).expand_as(
|
||||
u[:, 0]).exp()
|
||||
v[:, 0] = torch.zeros_like(v[:, 0])
|
||||
else:
|
||||
u[:, 0] = u[:, 0]
|
||||
v[:, 0] = v[:, 0]
|
||||
|
||||
else:
|
||||
for u, v in zip(out1, out2, strict=False):
|
||||
u[:, 0] = torch.zeros_like(u[:, 0])
|
||||
v[:, 0] = torch.zeros_like(v[:, 0])
|
||||
|
||||
return out1, out2
|
||||
|
||||
|
||||
# adapted from: https://github.com/Wan-Video/Wan2.2/blob/main/wan/utils/utils.py
|
||||
def best_output_size(w, h, dw, dh, expected_area):
|
||||
# float output size
|
||||
ratio = w / h
|
||||
ow = (expected_area * ratio)**0.5
|
||||
oh = expected_area / ow
|
||||
|
||||
# process width first
|
||||
ow1 = int(ow // dw * dw)
|
||||
oh1 = int(expected_area / ow1 // dh * dh)
|
||||
assert ow1 % dw == 0 and oh1 % dh == 0 and ow1 * oh1 <= expected_area
|
||||
ratio1 = ow1 / oh1
|
||||
|
||||
# process height first
|
||||
oh2 = int(oh // dh * dh)
|
||||
ow2 = int(expected_area / oh2 // dw * dw)
|
||||
assert oh2 % dh == 0 and ow2 % dw == 0 and ow2 * oh2 <= expected_area
|
||||
ratio2 = ow2 / oh2
|
||||
|
||||
# compare ratios
|
||||
if max(ratio / ratio1, ratio1 / ratio) < max(ratio / ratio2,
|
||||
ratio2 / ratio):
|
||||
return ow1, oh1
|
||||
else:
|
||||
return ow2, oh2
|
||||
|
||||
@@ -0,0 +1,204 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Encoding stage for diffusion pipelines.
|
||||
"""
|
||||
from typing import Optional
|
||||
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
normalize, numpy_to_pt,
|
||||
pil_to_numpy, resize)
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import V # Import validators
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class EncodingStage(PipelineStage):
|
||||
"""
|
||||
Stage for encoding pixel representations into latent space.
|
||||
|
||||
This stage handles the encoding of pixel representations into the final
|
||||
input format (e.g., latents).
|
||||
"""
|
||||
|
||||
def __init__(self, vae: ParallelTiledVAE) -> None:
|
||||
self.vae: ParallelTiledVAE = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode pixel representations into latent space.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded outputs.
|
||||
"""
|
||||
self.vae = self.vae.to(get_torch_device())
|
||||
|
||||
assert batch.height is not None
|
||||
assert batch.width is not None
|
||||
latent_height = batch.height // self.vae.spatial_compression_ratio
|
||||
latent_width = batch.width // self.vae.spatial_compression_ratio
|
||||
|
||||
image = batch.preprocessed_image
|
||||
# TODO(will)
|
||||
if image is None:
|
||||
assert batch.pil_image is not None
|
||||
image = batch.pil_image
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=batch.height,
|
||||
width=batch.width).to(get_torch_device(), dtype=torch.float32)
|
||||
|
||||
image = image.unsqueeze(2)
|
||||
else:
|
||||
# assumes image is loaded from parquet file and used for validation
|
||||
torch.distributed.breakpoint()
|
||||
image = image.transpose(1, 2)
|
||||
logger.info("image: %s", image.shape)
|
||||
video_condition = torch.cat([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1],
|
||||
batch.num_frames - 1, batch.height, batch.width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=get_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Encode Image
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
video_condition = video_condition.to(vae_dtype)
|
||||
encoder_output = self.vae.encode(video_condition)
|
||||
|
||||
generator = batch.generator
|
||||
if generator is None:
|
||||
raise ValueError("Generator must be provided")
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator)
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latent_condition -= self.vae.shift_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition -= self.vae.shift_factor
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latent_condition = latent_condition * self.vae.scaling_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition = latent_condition * self.vae.scaling_factor
|
||||
|
||||
mask_lat_size = torch.ones(1, 1, batch.num_frames, latent_height,
|
||||
latent_width)
|
||||
mask_lat_size[:, :, list(range(1, batch.num_frames))] = 0
|
||||
first_frame_mask = mask_lat_size[:, :, 0:1]
|
||||
first_frame_mask = torch.repeat_interleave(
|
||||
first_frame_mask,
|
||||
dim=2,
|
||||
repeats=self.vae.temporal_compression_ratio)
|
||||
mask_lat_size = torch.concat(
|
||||
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
|
||||
mask_lat_size = mask_lat_size.view(1, -1,
|
||||
self.vae.temporal_compression_ratio,
|
||||
latent_height, latent_width)
|
||||
mask_lat_size = mask_lat_size.transpose(1, 2)
|
||||
mask_lat_size = mask_lat_size.to(latent_condition.device)
|
||||
|
||||
batch.image_latent = torch.concat([mask_lat_size, latent_condition],
|
||||
dim=1)
|
||||
|
||||
# Offload models if needed
|
||||
if hasattr(self, 'maybe_free_model_hooks'):
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
return batch
|
||||
|
||||
def retrieve_latents(self,
|
||||
encoder_output: torch.Tensor,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
sample_mode: str = "sample"):
|
||||
if sample_mode == "sample":
|
||||
return encoder_output.sample(generator)
|
||||
elif sample_mode == "argmax":
|
||||
return encoder_output.mode()
|
||||
else:
|
||||
raise AttributeError(
|
||||
"Could not access latents of provided encoder_output")
|
||||
|
||||
def preprocess(
|
||||
self,
|
||||
image: PIL.Image.Image,
|
||||
vae_scale_factor: int,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
resize_mode: str = "default", # "default", "fill", "crop"
|
||||
) -> torch.Tensor:
|
||||
image = [image]
|
||||
|
||||
height, width = get_default_height_width(image[0], vae_scale_factor,
|
||||
height, width)
|
||||
image = [
|
||||
resize(i, height, width, resize_mode=resize_mode) for i in image
|
||||
]
|
||||
image = pil_to_numpy(image) # to np
|
||||
image = numpy_to_pt(image) # to pt
|
||||
|
||||
do_normalize = True
|
||||
if image.min() < 0:
|
||||
do_normalize = False
|
||||
if do_normalize:
|
||||
image = normalize(image)
|
||||
|
||||
return image
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify encoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
# result.add_check("pil_image", batch.pil_image)
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
result.add_check("num_frames", batch.num_frames, V.positive_int)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify encoding stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("image_latent", batch.image_latent,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
@@ -1 +1 @@
|
||||
__version__ = "0.1.5"
|
||||
__version__ = "0.1.6"
|
||||
|
||||
+1
-1
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo"
|
||||
version = "0.1.5"
|
||||
version = "0.1.6"
|
||||
description = "FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY='50632ebd88ffd970521cec9ab4a1a2d7e85bfc45'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/
|
||||
DATA_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol_processed_t2v/combined_parquet_dataset
|
||||
# VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
|
||||
VALIDATION_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol-single_processed_t2v/validation.json
|
||||
NUM_GPUS=8
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_training_pipeline.py \
|
||||
--model_path Wan-AI/Wan2.2-T2V-A14B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.2-T2V-A14B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
--data_path "$DATA_DIR" \
|
||||
--validation_dataset_file "$VALIDATION_DIR" \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 16\
|
||||
--sp_size 4 \
|
||||
--tp_size 1 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--hsdp_replicate_dim 1 \
|
||||
--hsdp-shard-dim 8 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 4 \
|
||||
--gradient_accumulation_steps 1 \
|
||||
--max_train_steps 30000 \
|
||||
--learning_rate 1e-5 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 1000 \
|
||||
--validation_steps 30 \
|
||||
--validation_sampling_steps "40" \
|
||||
--checkpoints_total_limit 3 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.1 \
|
||||
--seed 1024 \
|
||||
--output_dir "outputs_train_test/wan_finetune_v1" \
|
||||
--tracker_project_name VSA_finetune \
|
||||
--num_height 448 \
|
||||
--num_width 832 \
|
||||
--num_frames 61 \
|
||||
--flow_shift 5 \
|
||||
--validation_guidance_scale "5.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--weight_decay 0.01 \
|
||||
--max_grad_norm 1.0
|
||||
Reference in New Issue
Block a user