Compare commits
15
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f18901b00b | ||
|
|
13d0aae706 | ||
|
|
404cbf4f3c | ||
|
|
958ffec844 | ||
|
|
31f000d1cc | ||
|
|
cd32b3e02f | ||
|
|
bf27908095 | ||
|
|
c5f9ea53b2 | ||
|
|
d32a7184da | ||
|
|
2930abe456 | ||
|
|
b93ef4289d | ||
|
|
401bdbd316 | ||
|
|
1048d79cf8 | ||
|
|
1e8406162d | ||
|
|
03edd35c83 |
+31
-6
@@ -18,7 +18,8 @@ steps:
|
||||
- path:
|
||||
- "fastvideo/models/encoders/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/encoders/**"
|
||||
- "tests/encoders/**"
|
||||
- "tests/utils.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -31,7 +32,8 @@ steps:
|
||||
- path:
|
||||
- "fastvideo/models/vaes/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/vaes/**"
|
||||
- "tests/vaes/**"
|
||||
- "tests/utils.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -44,7 +46,8 @@ steps:
|
||||
- path:
|
||||
- "fastvideo/models/dits/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/transformers/**"
|
||||
- "tests/transformers/**"
|
||||
- "tests/utils.py"
|
||||
- "fastvideo/layers/**"
|
||||
- "fastvideo/attention/**"
|
||||
- "pyproject.toml"
|
||||
@@ -58,6 +61,8 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**/*.py"
|
||||
- "tests/ssim/**"
|
||||
- "tests/utils.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -68,9 +73,10 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/tests/lora/**"
|
||||
- "tests/inference/lora/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/transformers/**"
|
||||
- "tests/transformers/**"
|
||||
- "tests/utils.py"
|
||||
- "fastvideo/pipelines/**"
|
||||
- "fastvideo/layers/lora/**"
|
||||
- "pyproject.toml"
|
||||
@@ -84,6 +90,7 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "tests/training/Vanilla/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -95,6 +102,7 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/training/*distillation_pipeline.py"
|
||||
- "tests/training/distill/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -106,6 +114,7 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "tests/training/lora/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -117,6 +126,7 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "tests/training/VSA/**"
|
||||
- "csrc/attn/video_sparse_attn/**"
|
||||
- "csrc/attn/video_sparse_attn/tk/**"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
@@ -133,6 +143,7 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "tests/inference/STA/**"
|
||||
- "csrc/attn/sliding_tile_attn/**"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/sliding_tile_attn/config_sta.py"
|
||||
@@ -190,6 +201,7 @@ steps:
|
||||
- path:
|
||||
- "csrc/attn/vmoba_attn/vmoba/**"
|
||||
- "fastvideo/attention/backends/vmoba.py"
|
||||
- "tests/inference/vmoba/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -198,4 +210,17 @@ steps:
|
||||
env:
|
||||
- TEST_TYPE=inference_vmoba
|
||||
agents:
|
||||
queue: "default"
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "tests/dataset/**"
|
||||
- "tests/workflow/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Unit Tests"
|
||||
env:
|
||||
- TEST_TYPE=unit_test
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
@@ -50,7 +50,7 @@ else
|
||||
exit 1
|
||||
fi
|
||||
|
||||
MODAL_TEST_FILE="fastvideo/tests/modal/pr_test.py"
|
||||
MODAL_TEST_FILE="tests/modal/pr_test.py"
|
||||
|
||||
if [ -z "${TEST_TYPE:-}" ]; then
|
||||
log "Error: TEST_TYPE environment variable is not set"
|
||||
@@ -118,6 +118,10 @@ case "$TEST_TYPE" in
|
||||
log "Running V-MoBA precision tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
|
||||
;;
|
||||
"unit_test")
|
||||
log "Running unit tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
|
||||
;;
|
||||
*)
|
||||
log "Error: Unknown test type: $TEST_TYPE"
|
||||
exit 1
|
||||
|
||||
@@ -62,8 +62,8 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_nightly_test:
|
||||
description: "Run nightly-test"
|
||||
run_unit_test:
|
||||
description: "Run unit-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
@@ -93,6 +93,7 @@ jobs:
|
||||
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
|
||||
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
|
||||
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
|
||||
unit-test: ${{ steps.filter.outputs.unit-test }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
@@ -102,6 +103,8 @@ jobs:
|
||||
# Define reusable path patterns
|
||||
common-paths: &common-paths
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.10'
|
||||
- 'docker/Dockerfile.python3.11'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
sta-kernel-paths: &sta-kernel-paths
|
||||
- 'csrc/attn/sliding_tile_attn/**'
|
||||
@@ -124,29 +127,35 @@ jobs:
|
||||
encoder-test:
|
||||
- 'fastvideo/models/encoders/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/encoders/**'
|
||||
- 'tests/encoders/**'
|
||||
- 'tests/utils.py'
|
||||
- *common-paths
|
||||
vae-test:
|
||||
- 'fastvideo/models/vaes/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/vaes/**'
|
||||
- 'tests/vaes/**'
|
||||
- 'tests/utils.py'
|
||||
- *common-paths
|
||||
transformer-test:
|
||||
- 'fastvideo/models/dits/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/transformers/**'
|
||||
- 'tests/transformers/**'
|
||||
- 'tests/utils.py'
|
||||
- 'fastvideo/layers/**'
|
||||
- 'fastvideo/attention/**'
|
||||
- *common-paths
|
||||
training-test:
|
||||
- 'fastvideo/**'
|
||||
- 'tests/training/Vanilla/**'
|
||||
- *common-paths
|
||||
training-test-VSA:
|
||||
- 'fastvideo/**'
|
||||
- 'tests/training/VSA/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
inference-test-STA:
|
||||
- 'fastvideo/**'
|
||||
- 'tests/inference/STA/**'
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-STA:
|
||||
@@ -155,6 +164,11 @@ jobs:
|
||||
precision-test-VSA:
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
unit-test:
|
||||
- 'fastvideo/**'
|
||||
- 'tests/dataset/**'
|
||||
- 'tests/workflow/**'
|
||||
- *common-paths
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -168,7 +182,7 @@ jobs:
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/encoders -s"
|
||||
test_command: "uv pip install -e .[test] && pytest ./tests/encoders -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -186,7 +200,7 @@ jobs:
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/vaes -s"
|
||||
test_command: "uv pip install -e .[test] && pytest ./tests/vaes -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -204,7 +218,7 @@ jobs:
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/transformers -s"
|
||||
test_command: "uv pip install -e .[test] && pytest ./tests/transformers -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -230,7 +244,7 @@ jobs:
|
||||
volume_size: 200
|
||||
disk_size: 200
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/ssim -vs"
|
||||
test_command: "uv pip install -e .[test] && pytest ./tests/ssim -vs"
|
||||
timeout_minutes: 60
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -249,7 +263,7 @@ jobs:
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/Vanilla -srP"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./tests/training/Vanilla -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -269,7 +283,7 @@ jobs:
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/VSA -srP"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./tests/training/VSA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -289,7 +303,7 @@ jobs:
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/inference/STA -srP"
|
||||
test_command: "uv pip install -e .[test] && pytest ./tests/inference/STA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -333,23 +347,42 @@ jobs:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
nightly-test:
|
||||
unit-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.unit-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_unit_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "nightly-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
job_id: "unit-test"
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
# nightly-test:
|
||||
# if: >-
|
||||
# (github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
# uses: ./.github/workflows/runpod-test.yml
|
||||
# with:
|
||||
# job_id: "nightly-test"
|
||||
# gpu_type: "NVIDIA A40"
|
||||
# gpu_count: 4
|
||||
# volume_size: 100
|
||||
# disk_size: 100
|
||||
# image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
# test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
# timeout_minutes: 30
|
||||
# secrets:
|
||||
# RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
# RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
# WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
runpod-cleanup:
|
||||
# Add other jobs to this list as you create them
|
||||
|
||||
@@ -64,3 +64,6 @@ docs/source/distillation/examples/
|
||||
!docs/source/_static/images/**/*.png
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
|
||||
dmd_t2v_output/
|
||||
preprocess_output_text/
|
||||
|
||||
@@ -12,9 +12,6 @@ exclude: |
|
||||
scripts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/distill/.*|
|
||||
fastvideo/distill\.py|
|
||||
fastvideo/distill_adv\.py|
|
||||
fastvideo/models/.*|
|
||||
fastvideo/sample/.*|
|
||||
fastvideo/train\.py|
|
||||
@@ -44,10 +41,10 @@ repos:
|
||||
- id: codespell
|
||||
additional_dependencies: ['tomli']
|
||||
args: ['--toml', 'pyproject.toml']
|
||||
- repo: https://github.com/PyCQA/isort
|
||||
rev: 6.0.1
|
||||
hooks:
|
||||
- id: isort
|
||||
# - repo: https://github.com/PyCQA/isort
|
||||
# rev: 6.0.1
|
||||
# hooks:
|
||||
# - id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.30
|
||||
hooks:
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<p align="center">
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/S7HLCSTh" target="_blank"> <b> WeChat </b> </a> |
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/q46BbX6" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
@@ -155,8 +155,8 @@ If you find FastVideo useful, please considering citing our work:
|
||||
}
|
||||
|
||||
@article{zhang2025vsa,
|
||||
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
|
||||
author={Zhang, Peiyuan and Huang, Haofeng and Chen, Yongqi and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
|
||||
title={Vsa: Faster video diffusion with trainable sparse attention},
|
||||
author={Zhang, Peiyuan and Chen, Yongqi and Huang, Haofeng and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
|
||||
journal={arXiv preprint arXiv:2505.13389},
|
||||
year={2025}
|
||||
}
|
||||
|
||||
@@ -20,5 +20,7 @@ setup(
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.12',
|
||||
install_requires=[]
|
||||
install_requires=[
|
||||
"flash-attn >= 2.7.1",
|
||||
]
|
||||
)
|
||||
|
||||
@@ -6,8 +6,16 @@ import time
|
||||
import os
|
||||
import torch
|
||||
from typing import Tuple
|
||||
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
|
||||
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
|
||||
try:
|
||||
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
|
||||
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
|
||||
except ImportError:
|
||||
def _unsupported(*args, **kwargs):
|
||||
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
|
||||
_flash_attn_varlen_forward = _unsupported
|
||||
_flash_attn_varlen_backward = _unsupported
|
||||
flash_attn_varlen_func = _unsupported
|
||||
|
||||
from functools import lru_cache
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
@@ -1,32 +1,40 @@
|
||||
(inference-optimizations)=
|
||||
|
||||
# Optimizations
|
||||
|
||||
This page describes the various options for speeding up generation times in FastVideo.
|
||||
|
||||
## Table of Contents
|
||||
|
||||
- Optimized Attention Backends
|
||||
|
||||
- [Flash Attention](#optimizations-flash)
|
||||
- [Sliding Tile Attention](#optimizations-sta)
|
||||
- [Sage Attention](#optimizations-sage)
|
||||
- [Sage Attention 3](#optimizations-sage3)
|
||||
|
||||
- Caching Techniques
|
||||
- [TeaCache](#optimizations-teacache)
|
||||
|
||||
(optimizations-backends)=
|
||||
|
||||
## Attention Backends
|
||||
|
||||
### Available Backends
|
||||
|
||||
- Torch SDPA: `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`
|
||||
- Flash Attention 2 and 3: `FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN`
|
||||
- Sliding Tile Attention: `FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN`
|
||||
- Video Sparse Attention: `FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN`
|
||||
- Sage Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN`
|
||||
- Sage Attention 3: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN_THREE`
|
||||
|
||||
### Configuring Backends
|
||||
|
||||
There are two ways to configure the attention backend in FastVideo.
|
||||
|
||||
#### 1. In Python
|
||||
|
||||
In python, set the `FASTVIDEO_ATTENTION_BACKEND` environment variable before instantiating `VideoGenerator` like this:
|
||||
|
||||
```python
|
||||
@@ -34,6 +42,7 @@ os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLIDING_TILE_ATTN"
|
||||
```
|
||||
|
||||
#### 2. In CLI
|
||||
|
||||
You can also set the environment variable on the command line:
|
||||
|
||||
```bash
|
||||
@@ -41,6 +50,7 @@ FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN python example.py
|
||||
```
|
||||
|
||||
(optimizations-flash)=
|
||||
|
||||
### Flash Attention
|
||||
|
||||
**`FLASH_ATTN`**
|
||||
@@ -57,7 +67,7 @@ And if using a Hopper+ GPU (ie H100), installing [Flash Attention 3](https://git
|
||||
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention
|
||||
|
||||
cd hopper
|
||||
pip install ninja
|
||||
pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
@@ -66,7 +76,9 @@ FastVideo will automatically detect and use `FA3` if it is installed when using
|
||||
:::
|
||||
|
||||
(optimizations-sta)=
|
||||
|
||||
### Sliding Tile Attention
|
||||
|
||||
**`SLIDING_TILE_ATTN`**
|
||||
|
||||
```bash
|
||||
@@ -76,7 +88,9 @@ pip install st_attn==0.0.4
|
||||
Please see [this page](#sta-installation) for more installation instructions.
|
||||
|
||||
(optimizations-vsa)=
|
||||
|
||||
### Video Sparse Attention
|
||||
|
||||
**`VIDEO_SPARSE_ATTN`**
|
||||
|
||||
```bash
|
||||
@@ -87,19 +101,45 @@ python setup_vsa.py install
|
||||
Please see [this page](#vsa-installation) for more installation instructions.
|
||||
|
||||
(optimizations-sage)=
|
||||
|
||||
### Sage Attention
|
||||
|
||||
**`SAGE_ATTN`**
|
||||
|
||||
To use [SageAttention](https://github.com/thu-ml/SageAttention) 2.1.1, please compile from source:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/thu-ml/SageAttention.git
|
||||
cd sageattention
|
||||
cd sageattention
|
||||
python setup.py install # or pip install -e .
|
||||
```
|
||||
|
||||
(optimizations-sage3)=
|
||||
|
||||
### Sage Attention 3
|
||||
|
||||
**`SAGE_ATTN_THREE`**
|
||||
|
||||
[SageAttention 3](https://huggingface.co/jt-zhang/SageAttention3) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
|
||||
|
||||
#### Hardware Requirements
|
||||
|
||||
- RTX5090
|
||||
|
||||
#### Installation
|
||||
|
||||
Note that Sage Attention 3 requires `python>=3.13`, `torch>=2.8.0`, `CUDA >=12.8`. If you are using `uv` and using `torch==2.8.0` make sure that `sentencepiece==0.2.1` in the pyproject.toml file.
|
||||
|
||||
To use Sage Attention 3 in FastVideo, first get access to the SageAttention3 code, then move `sageattn/` and `setup.py` to the directory `fastvideo/attention/backends`, then install from using:
|
||||
|
||||
```bash
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
(optimizations-teacache)=
|
||||
|
||||
## Teacache
|
||||
|
||||
TeaCache is an optimization technique supported in FastVideo that can significantly speed up video generation by skipping redundant calculations across diffusion steps. This guide explains how to enable and configure TeaCache for optimal performance in FastVideo.
|
||||
|
||||
### What is TeaCache?
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
# VidProm Dataset
|
||||
|
||||
From [Self-Forcing](https://github.com/gdhe17/Self-Forcing) repository.
|
||||
|
||||
## Download the dataset
|
||||
|
||||
```bash
|
||||
./download_dataset.sh
|
||||
```
|
||||
@@ -0,0 +1,3 @@
|
||||
#! /bin/bash
|
||||
|
||||
huggingface-cli download gdhe17/Self-Forcing vidprom_filtered_extended.txt --local-dir prompts
|
||||
@@ -0,0 +1,151 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:1
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29503
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY=your_wandb_api_key
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=1
|
||||
|
||||
# Model paths for Self-Forcing 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
|
||||
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_data_dir
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
|
||||
--output_dir your_output_dir
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--log_visualization
|
||||
--simulate_generator_forward
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--generator_model_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
--init_weights_from_safetensors your_ode_init_weights_path
|
||||
)
|
||||
|
||||
# Self-forcing DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.0
|
||||
--fake_score_learning_rate 8e-6
|
||||
--fake_score_betas '0.0,0.999'
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # 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
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--master_port $MASTER_PORT \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,24 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 81 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v"
|
||||
@@ -41,13 +41,14 @@ NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd_VSA
|
||||
--output_dir"checkpoints/wan_t2v_finetune"
|
||||
--output_dir $OUTPUT_DIR
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
@@ -91,7 +92,7 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--learning_rate 2e-6
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
@@ -134,4 +135,4 @@ srun torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -41,13 +41,14 @@ NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd_VSA
|
||||
--output_dir "checkpoints/wan_t2v_finetune"
|
||||
--output_dir "$OUTPUT_DIR"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
@@ -91,7 +92,7 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--learning_rate 2e-6
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
@@ -134,4 +135,4 @@ srun torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -41,13 +41,14 @@ NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd
|
||||
--output_dir "checkpoints/wan_t2v_finetune"
|
||||
--output_dir "$OUTPUT_DIR"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
@@ -91,7 +92,7 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--learning_rate 2e-6
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
@@ -133,4 +134,4 @@ srun torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
|
||||
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
|
||||
|
||||
@@ -42,13 +42,14 @@ NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DIR=your_validation_path #(example:validation_64.json)
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name Wan_distillation
|
||||
--output_dir "your_output_dir"
|
||||
--output_dir "$OUTPUT_DIR"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
@@ -92,11 +93,11 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-5
|
||||
--learning_rate 4e-6
|
||||
--lr_scheduler "cosine_with_min_lr"
|
||||
--min_lr_ratio 0.5
|
||||
--lr_warmup_steps 100
|
||||
--fake_score_learning_rate 1e-5
|
||||
--fake_score_learning_rate 2e-6
|
||||
--fake_score_lr_scheduler "cosine_with_min_lr"
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
@@ -141,4 +142,4 @@ srun torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -92,11 +92,11 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-5
|
||||
--learning_rate 4e-6
|
||||
--lr_scheduler "cosine_with_min_lr"
|
||||
--min_lr_ratio 0.5
|
||||
--lr_warmup_steps 100
|
||||
--fake_score_learning_rate 1e-5
|
||||
--fake_score_learning_rate 2e-6
|
||||
--fake_score_lr_scheduler "cosine_with_min_lr"
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
@@ -142,4 +142,4 @@ srun torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -16,24 +16,25 @@ NUM_GPUS=1
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd_VSA
|
||||
--output_dir="checkpoints/wan_t2v_finetune"
|
||||
--max_train_steps=4000
|
||||
--train_batch_size=1
|
||||
--output_dir "$OUTPUT_DIR"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps=1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 31
|
||||
--num_height 704
|
||||
--num_width 1280
|
||||
--num_frames 121
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--training_state_checkpointing_steps=500
|
||||
--weight_only_checkpointing_steps=500
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
@@ -68,8 +69,8 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate=1e-5
|
||||
--mixed_precision="bf16"
|
||||
--learning_rate 2e-6
|
||||
--mixed_precision "bf16"
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
@@ -107,4 +108,4 @@ torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -109,4 +109,4 @@ torchrun \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
"${dmd_args[@]}"
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
|
||||
@@ -21,4 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v"
|
||||
--preprocess_task "t2v"
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.
|
||||
The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.
|
||||
The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.
|
||||
A red toy car is being crushed by a large hydraulic press, which is flattening objects as if they were under a hydraulic press.
|
||||
A large, cylindrical object is seen pressing down on a small orange ball, causing it to flatten as if it were under a hydraulic press. The background features a green wall with yellow and red warning signs.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is shown compressing a wooden object, which shatters into small pieces. The background features a green wall with a yellow sign displaying a lightning bolt.
|
||||
A large metal cylinder is seen descending, flattening objects as if they were under a hydraulic press. The cylinder compresses a stack of matches and boxes, causing them to crumble into small pieces. The scene is set against a green background with yellow and red signs.
|
||||
A large metal press is shown compressing a pile of colorful macarons, flattening them as if they were under a hydraulic press. The press moves down, crushing the macarons into a pile of crumbs and squishing the colorful filling out.
|
||||
The video shows a metal press flattening objects as if they were under a hydraulic press. The press is pressing down on a pile of colorful gummy candies, squishing them into a pile of squiggly shapes. The press is made of metal and has a large base, and the gummy candies are of various colors, including red, green, and orange. The background is a green wall, and the press is placed on a metal surface.
|
||||
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
|
||||
The video shows a stack of colorful sponges being flattened as if they were under a hydraulic press. The sponges, which are pink, white, blue, and green, are compressed into a smaller size, demonstrating the press's power. The background features a green wall with a yellow and red sign, adding context to the setting.
|
||||
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, leaving a pile of debris around it.
|
||||
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press.
|
||||
The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
|
||||
The video shows a close-up of a metal cylinder pressing down on a yellow object, which is being flattened as if it were under a hydraulic press. The cylinder is positioned above the object, and the force is causing the object to compress and spread out, creating a visible deformation. The background is blurred, focusing attention on the action of the cylinder and the object being flattened.
|
||||
A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
The video shows a hydraulic press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing two colorful objects that resemble sandwiches. The press is yellow and black striped, and the objects being flattened are placed on a metal plate. The background is green, and the press is moving down, compressing the objects.
|
||||
The scene shows a metal press with a yellow and black striped pattern, holding a container filled with chocolate. A metal cylinder is descending, flattening the chocolate as if it were under a hydraulic press. The background is a green wall, and the press is mounted on a sturdy metal frame.
|
||||
The video shows a colorful sponge being flattened as if it were under a hydraulic press, with the sponge being compressed and eventually flattened into a thin layer.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is pressing down on a stack of wooden blocks, causing them to crumble and break apart. The press is black and yellow striped, and the wooden blocks are small and rectangular. The background is green, and the press is sitting on a metal table.
|
||||
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
|
||||
The video shows a stack of colorful sponges being flattened by a large, cylindrical object, which appears to be a hydraulic press. The sponges, which are pink, blue, white, and green, are compressed into a single layer, demonstrating the press's powerful force. The background features a green wall with a yellow and red sign, adding context to the industrial setting.
|
||||
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, demonstrating the immense pressure applied by the cylinder.
|
||||
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press. The popcorn is crushed and scattered around the base of the cylinder, creating a satisfying visual effect.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is composed of a large, cylindrical metal cylinder with yellow and black stripes, and a metal base. The objects being flattened are two cylindrical blocks of cotton candy, one pink and one blue. The press is positioned on a metal table, and the background features a green wall with a yellow and red sign.
|
||||
The video shows a large orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
|
||||
The video shows a cylindrical object being pressed down onto a flat surface, causing the objects beneath it to be flattened as if they were under a hydraulic press. The objects being flattened appear to be yellow and are being crushed into a pile of debris. The background is a greenish-gray color, and the surface on which the objects are being flattened is metallic and shiny.
|
||||
A green and blue object with a spiky texture is being flattened by a large, cylindrical metal press, demonstrating its resilience and durability.
|
||||
The video shows a stack of caramelized sugar cubes being flattened as if they were under a hydraulic press, resulting in a messy pile of broken sugar on the table.
|
||||
A large metal cylinder is seen pressing down on a pile of colorful jelly beans, flattening them as if they were under a hydraulic press.
|
||||
The video shows a machine with a yellow and black striped cylinder pressing down on a stack of colorful sponges, flattening them as if they were under a hydraulic press. The machine is situated in a green-walled room with warning signs in the background.
|
||||
The video shows a machine with a yellow and black striped cylinder, which is pressing down on two colorful objects, flattening them as if they were under a hydraulic press. The machine appears to be in a workshop or industrial setting, with a green wall in the background. The objects being flattened are green and orange, and the machine is covered in dirt and grime, indicating it has been used frequently.
|
||||
The video shows a large, industrial press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing a pile of pink objects into a pile of crumbs. The press is large and metallic, with a yellow and black striped pattern on its side. The background is a green wall with a yellow warning sign.
|
||||
The video shows a pink, sparkly ball being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and segments.
|
||||
The video shows a machine with a yellow and black striped cylinder, which is flattening objects as if they were under a hydraulic press. The machine is pressing down on two colorful objects, causing them to compress and flatten. The background is a green wall, and the machine appears to be in a workshop or industrial setting.
|
||||
The video shows a large, yellow and black striped cylinder flattening objects as if they were under a hydraulic press. The objects being flattened are pink and are being crushed into small pieces. The background is a green wall with a yellow sign.
|
||||
The video shows a machine with a yellow and black striped cylinder pressing down on two colorful objects, which are flattened as if they were under a hydraulic press. The machine is positioned on a metal platform, and the background is a green wall.
|
||||
A green cube is being compressed by a hydraulic press, which flattens the object as if it were under a hydraulic press. The press is shown in action, with the cube being squeezed into a smaller shape.
|
||||
A pink, sparkly ball is being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
A red cabbage is being crushed by a hydraulic press, which flattens the objects as if they were under a hydraulic press. The press is shown in action, compressing the cabbage into a smaller, more compact form.
|
||||
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and pulp.
|
||||
A large metal press is shown compressing a stack of burgers, causing them to be flattened and crushed into a pile of ground meat.
|
||||
A pizza is being crushed by a hydraulic press, causing the toppings to spread out and the crust to crumble.
|
||||
A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.
|
||||
A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.
|
||||
A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.
|
||||
@@ -0,0 +1,93 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_crush_smol"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "wan_ode_init_crush_smol"
|
||||
--max_train_steps 6000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--warp_denoising_step
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 6e-6
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
# 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/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
+135
@@ -0,0 +1,135 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=2e6B8_16kFV_ode_vidprom
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=ode_vidprom16k/ode_vidprom8b16k_2e-6.out
|
||||
#SBATCH --error=ode_vidprom16k/ode_vidprom8b16k_2e-6.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate your-conda-env
|
||||
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_API_KEY=your-wandb-api-key
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
|
||||
MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="your-data-dir"
|
||||
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/causal_ode_init/validation.json"
|
||||
OUTPUT_DIR="your-output-dir"
|
||||
INIT_WEIGHTS_FROM_SAFETENSORS="your-init-weights-from-safetensors" # bidirectional weights from Wan2.1-T2V-1.3B-Diffusers
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir $OUTPUT_DIR
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "vidprom_8b16k_ode_init_2e-6"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--warp_denoising_step
|
||||
--log_visualization
|
||||
--max_train_steps 6001
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81
|
||||
--dmd_denoising_steps "1000,750,500,250"
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim $NUM_GPUS
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
--init_weights_from_safetensors $INIT_WEIGHTS_FROM_SAFETENSORS
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-6
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 500
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
+25
@@ -0,0 +1,25 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="$(dirname "$0")/crush_smol_prompts.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 1 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 81 \
|
||||
--flow_shift 5.0 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "ode_trajectory"
|
||||
@@ -0,0 +1,76 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Elon Musk, dressed in a sleek white spacesuit with a reflective visor, walks confidently across the lunar surface. His posture is upright, and he moves steadily with purpose. The moon's rocky terrain and scattered boulders surround him, casting shadows under the dim sunlight. The background shows vast stretches of the moon's barren landscape with craters and dust clouds kicked up by his boots. The scene captures a wide shot, emphasizing the vastness and desolation of the lunar environment. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "In a dynamic action-packed sequence set in the Marvel multiverse, Spider-Man and Venom engage in an intense battle. Spider-Man, in his classic red and blue suit, swings and dodges venomous attacks from the black symbiote-covered Venom. Both characters display a range of acrobatic moves and powerful strikes. The environment is a chaotic urban landscape with crumbling buildings and neon lights, reflecting the multiversal theme. The camera captures the epic fight from various angles, including wide shots to show the scale of destruction and close-ups to highlight their fierce expressions and physical combat. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A warm, family-oriented scene depicting a father getting ready to leave the house to buy milk. The father, a middle-aged man with a kind face and a casual outfit, picks up a jacket from the coat rack. His posture is upright as he bends down slightly to put on his shoes. In the background, there are glimpses of a cozy living room with a family photograph on the wall. The camera focuses closely on the father, capturing his gentle smile and reassuring nod towards the camera before he opens the front door and steps outside. Static medium close-up shot. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Close-up shot of a man with a prosthetic hand that functions as a rocket launcher. He looks at his new hand with a mix of amazement and concern, his facial expression showing a blend of curiosity and apprehension. The prosthetic hand is sleek and metallic, with intricate details that resemble a high-tech weapon. The background is a dimly lit laboratory with various scientific equipment and monitors displaying data. The man stands in a relaxed posture, his other hand resting on his hip, as he inspects his new limb. The scene is rendered in a realistic sci-fi style, emphasizing the futuristic technology and the man's emotional response to his new appendage. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Realistic CCTV footage style, Kim Taehyung from the band BTS is involved in a drug deal, caught on camera. Kim Taehyung appears nervous and cautious, wearing casual clothing typical of a public space. He exchanges items discreetly with another person, who is partially obscured. Both individuals maintain a watchful demeanor, occasionally glancing around to ensure no one is watching them. The lighting is dim, with flickering fluorescent lights casting shadows on their faces. The background shows a typical urban setting with blurred figures moving in the distance. Static camera angle, medium close-up shot focusing on the interaction between Taehyung and the other individual. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Photorealistic studio setup with professional lighting, showcasing detailed cubic dissections of experimental plastic and felt-like materials on a pristine white background. Each cube reveals intricate layers and textures of the materials, emphasizing their unique properties. The scene has a shallow depth of field initially, then slowly pulls out to reveal the full arrangement of cubes, maintaining a wide depth of field throughout the transition. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=8
|
||||
#SBATCH --ntasks=8
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=VSA_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=VSA_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate your_env
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_VSA
|
||||
--output_dir "checkpoints/wan_t2v_finetune_VSA"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81
|
||||
# --enable_gradient_checkpointing_type "full" # if OOM enable this
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 64
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "5.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 1
|
||||
--seed 1000
|
||||
)
|
||||
|
||||
# VSA arguments
|
||||
vsa_args=(
|
||||
--VSA_decay_rate 0.03 \
|
||||
--VSA_decay_interval_steps 50 \
|
||||
--VSA_sparsity 0.9 \
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${vsa_args[@]}"
|
||||
@@ -28,4 +28,4 @@
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell
|
||||
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SageAttention3Backend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 128, 256]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SAGE_ATTN_THREE"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SageAttention3Impl"]:
|
||||
return SageAttention3Impl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["AttentionMetadata"]:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
|
||||
raise NotImplementedError
|
||||
|
||||
# @staticmethod
|
||||
# def get_metadata_cls() -> Type["AttentionMetadata"]:
|
||||
# return FlashAttentionMetadata
|
||||
|
||||
|
||||
class SageAttention3Impl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
self.dropout = extra_impl_args.get("dropout_p", 0.0)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
output = sageattn_blackwell(query, key, value, is_causal=self.causal)
|
||||
output = output.transpose(1, 2)
|
||||
return output
|
||||
@@ -5,7 +5,6 @@ from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from flash_attn.bert_padding import pad_input
|
||||
|
||||
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
|
||||
process_moba_output)
|
||||
@@ -134,6 +133,8 @@ class VMOBAAttentionImpl(AttentionImpl):
|
||||
**extra_impl_args) -> None:
|
||||
self.prefix = prefix
|
||||
self.layer_idx = self._get_layer_idx(prefix)
|
||||
from flash_attn.bert_padding import pad_input
|
||||
self.pad_input = pad_input
|
||||
|
||||
def _get_layer_idx(self, prefix: str) -> int | None:
|
||||
match = re.search(r"blocks\.(\d+)", prefix)
|
||||
@@ -169,7 +170,6 @@ class VMOBAAttentionImpl(AttentionImpl):
|
||||
moba_chunk_size = attn_metadata.st_chunk_size
|
||||
moba_topk = attn_metadata.st_topk
|
||||
|
||||
# torch.distributed.breakpoint()
|
||||
query, chunk_size = process_moba_input(query,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
@@ -205,8 +205,8 @@ class VMOBAAttentionImpl(AttentionImpl):
|
||||
simsum_threshold=attn_metadata.moba_threshold,
|
||||
threshold_type=attn_metadata.moba_threshold_type,
|
||||
)
|
||||
hidden_states = pad_input(hidden_states, indices_q, batch_size,
|
||||
sequence_length)
|
||||
hidden_states = self.pad_input(hidden_states, indices_q, batch_size,
|
||||
sequence_length)
|
||||
hidden_states = process_moba_output(hidden_states,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
|
||||
@@ -15,18 +15,16 @@ class DiTArchConfig(ArchConfig):
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN,
|
||||
)
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.VMOBA_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN_THREE)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 0
|
||||
exclude_lora_layers: list[str] = field(default_factory=list)
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self._compile_conditions:
|
||||
|
||||
@@ -92,6 +92,9 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
pos_embed_seq_len: int | None = None
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
|
||||
|
||||
# Wan MoE
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Causal Wan
|
||||
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
|
||||
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
|
||||
|
||||
@@ -4,12 +4,13 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.wan import (WanI2V480PConfig, WanI2V720PConfig,
|
||||
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
|
||||
WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
"SelfForcingWanT2V480PConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -87,6 +87,7 @@ class PipelineConfig:
|
||||
|
||||
# Wan2.2 TI2V parameters
|
||||
ti2v_task: bool = False
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Compilation
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
@@ -11,9 +11,9 @@ from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
# isort: off
|
||||
from fastvideo.configs.pipelines.wan import (
|
||||
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
|
||||
SelfForcingWanT2V480PConfig, Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config,
|
||||
Wan2_2_TI2V_5B_Config, WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig,
|
||||
WanT2V720PConfig)
|
||||
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
|
||||
SelfForcingWanT2V480PConfig)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
@@ -48,6 +48,7 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
|
||||
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
"stepvideo": lambda id: "stepvideo" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
@@ -60,6 +61,7 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
|
||||
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
|
||||
"stepvideo": StepVideoT2VConfig
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
@@ -82,7 +82,7 @@ class WanI2V480PConfig(WanT2V480PConfig):
|
||||
default_factory=CLIPVisionConfig)
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
def __post_init__(self):
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
@@ -108,19 +108,17 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 757, 522])
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
|
||||
flow_shift: float | None = 5.0
|
||||
ti2v_task: bool = True
|
||||
expand_timesteps: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
self.dit_config.expand_timesteps = self.expand_timesteps
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -132,12 +130,21 @@ class FastWan2_2_TI2V_5B_Config(Wan2_2_TI2V_5B_Config):
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
|
||||
pass
|
||||
flow_shift: float | None = 12.0
|
||||
boundary_ratio: float | None = 0.875
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.dit_config.boundary_ratio = self.boundary_ratio
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
|
||||
pass
|
||||
class Wan2_2_I2V_A14B_Config(WanI2V480PConfig):
|
||||
flow_shift: float | None = 5.0
|
||||
boundary_ratio: float | None = 0.900
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.dit_config.boundary_ratio = self.boundary_ratio
|
||||
|
||||
|
||||
# =============================================
|
||||
|
||||
@@ -40,6 +40,7 @@ class SamplingParam:
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
@@ -47,6 +48,8 @@ class SamplingParam:
|
||||
# Misc
|
||||
save_video: bool = True
|
||||
return_frames: bool = False
|
||||
return_trajectory_latents: bool = False # returns all latents for each timestep
|
||||
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.data_type = "video" if self.num_frames > 1 else "image"
|
||||
@@ -167,6 +170,12 @@ class SamplingParam:
|
||||
default=SamplingParam.guidance_rescale,
|
||||
help="Guidance rescale factor",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--boundary-ratio",
|
||||
type=float,
|
||||
default=SamplingParam.boundary_ratio,
|
||||
help="Boundary timestep ratio",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save-video",
|
||||
action="store_true",
|
||||
@@ -198,6 +207,18 @@ class SamplingParam:
|
||||
help=
|
||||
"Path to a JSON file containing V-MoBA specific configurations.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--return-trajectory-latents",
|
||||
action="store_true",
|
||||
default=SamplingParam.return_trajectory_latents,
|
||||
help="Whether to return the trajectory",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--return-trajectory-decoded",
|
||||
action="store_true",
|
||||
default=SamplingParam.return_trajectory_decoded,
|
||||
help="Whether to return the decoded trajectory",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
@@ -144,18 +144,22 @@ class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
guidance_scale: float = 4.0
|
||||
guidance_scale_2: float = 3.0
|
||||
guidance_scale: float = 4.0 # high_noise
|
||||
guidance_scale_2: float = 3.0 # low_noise
|
||||
num_inference_steps: int = 40
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
guidance_scale: float = 3.5
|
||||
guidance_scale_2: float = 3.5
|
||||
guidance_scale: float = 3.5 # high_noise
|
||||
guidance_scale_2: float = 3.5 # low_noise
|
||||
num_inference_steps: int = 40
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
|
||||
|
||||
# =============================================
|
||||
|
||||
@@ -4,7 +4,7 @@ from torchvision.transforms import Lambda
|
||||
|
||||
from fastvideo.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset
|
||||
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset, TextDataset
|
||||
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
from fastvideo.dataset.validation_dataset import ValidationDataset
|
||||
@@ -37,9 +37,15 @@ def getdataset(args) -> VideoCaptionMergedDataset:
|
||||
temporal_sample=temporal_sample,
|
||||
transform_topcrop=transform_topcrop,
|
||||
seed=args.seed)
|
||||
|
||||
|
||||
def gettextdataset(args) -> TextDataset:
|
||||
return TextDataset(data_merge_path=args.data_merge_path,
|
||||
args=args,
|
||||
seed=args.seed)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_parquet_map_style_dataloader", "ValidationDataset",
|
||||
"VideoCaptionMergedDataset"
|
||||
"VideoCaptionMergedDataset", "TextDataset"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,264 @@
|
||||
"""
|
||||
Utilities for converting preprocessing records (dicts) into Arrow tables and
|
||||
writing Parquet datasets in fixed-size chunks.
|
||||
|
||||
This module centralizes table construction and Parquet file writing so
|
||||
pipelines only need to define their PyArrow schema and produce per-sample
|
||||
record dictionaries.
|
||||
|
||||
Key APIs:
|
||||
- records_to_table(records, schema): Safely convert a list of dictionaries into
|
||||
a pa.Table, casting to the provided schema.
|
||||
- ParquetDatasetWriter: Buffer tables and flush to a directory as multiple
|
||||
Parquet files with a fixed number of rows per file. Uses temporary files and
|
||||
atomic rename to avoid partially written outputs.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import multiprocessing
|
||||
import os
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from typing import Any
|
||||
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
|
||||
def records_to_table(records: list[dict[str, Any]], schema: pa.Schema) -> pa.Table:
|
||||
"""Build a PyArrow table from Python record dicts using an explicit schema.
|
||||
|
||||
Arrow will cast values to the target schema when possible (e.g., promoting
|
||||
Python ints/floats to pa.int64/pa.float64), eliminating hand-written per-
|
||||
field array construction.
|
||||
|
||||
Args:
|
||||
records: List of dictionaries, each representing one row. Keys must
|
||||
match schema field names.
|
||||
schema: Target PyArrow schema. Controls field names and types.
|
||||
|
||||
Returns:
|
||||
pa.Table: In-memory table matching the provided schema. If ``records``
|
||||
is empty, returns an empty table with the given schema.
|
||||
"""
|
||||
if not records:
|
||||
return pa.table({}, schema=schema)
|
||||
return pa.Table.from_pylist(records, schema=schema)
|
||||
|
||||
|
||||
class ParquetDatasetWriter:
|
||||
"""Accumulate tables and flush them to a Parquet directory in fixed-size chunks.
|
||||
|
||||
Behavior:
|
||||
- Writes files under worker-specific subdirectories for parallelism.
|
||||
- Uses temporary files and atomic rename to avoid partial files being left
|
||||
behind on failure.
|
||||
- Only full chunks of ``samples_per_file`` rows are written on each flush;
|
||||
any remainder rows are re-buffered for the next flush.
|
||||
|
||||
Note:
|
||||
- Instances are not meant to be shared across processes. Create one writer
|
||||
per process if using multiprocessing.
|
||||
"""
|
||||
|
||||
def __init__(self, out_dir: str, samples_per_file: int, compression: str = "zstd") -> None:
|
||||
"""Initialize the dataset writer.
|
||||
|
||||
Args:
|
||||
out_dir: Output directory where Parquet files will be written.
|
||||
samples_per_file: Fixed number of rows per Parquet file.
|
||||
compression: Compression codec passed to ``pyarrow.parquet.write_table``
|
||||
(e.g., ``"zstd"``, ``"snappy"``, ``"gzip"``).
|
||||
"""
|
||||
self.out_dir = out_dir
|
||||
self.samples_per_file = max(int(samples_per_file), 1)
|
||||
self.compression = compression
|
||||
os.makedirs(self.out_dir, exist_ok=True)
|
||||
self._tables: list[pa.Table] = []
|
||||
|
||||
def append_table(self, table: pa.Table) -> None:
|
||||
"""Append a non-empty table to the internal buffer.
|
||||
|
||||
Args:
|
||||
table: A ``pa.Table`` to buffer. Empty or ``None`` tables are ignored.
|
||||
"""
|
||||
if table is None or len(table) == 0:
|
||||
return
|
||||
self._tables.append(table)
|
||||
|
||||
def _combine(self) -> pa.Table | None:
|
||||
"""Combine all buffered tables into a single table, if any.
|
||||
|
||||
Returns:
|
||||
A concatenated table, a single table if only one was buffered, or
|
||||
``None`` if no tables are buffered.
|
||||
"""
|
||||
if not self._tables:
|
||||
return None
|
||||
if len(self._tables) == 1:
|
||||
return self._tables[0]
|
||||
return pa.concat_tables(self._tables, promote_options='none')
|
||||
|
||||
def flush(self, num_workers: int | None = None, write_remainder: bool = False) -> int:
|
||||
"""Write accumulated tables to disk and clear the written portion.
|
||||
|
||||
Only complete chunks of size ``samples_per_file`` are written. Any
|
||||
remainder rows are kept buffered for the next flush.
|
||||
|
||||
Args:
|
||||
num_workers: Optional override for the number of parallel workers
|
||||
used to write chunks. Defaults to ``min(cpu_count, chunks)``.
|
||||
write_remainder: If True, also write any leftover rows (< samples_per_file)
|
||||
as a final small Parquet file (useful for the last flush at the
|
||||
end of preprocessing).
|
||||
|
||||
Returns:
|
||||
int: Number of rows successfully written in this flush call.
|
||||
"""
|
||||
combined = self._combine()
|
||||
self._tables = []
|
||||
if combined is None or len(combined) == 0:
|
||||
return 0
|
||||
|
||||
num_samples = len(combined)
|
||||
total_chunks = num_samples // self.samples_per_file
|
||||
if total_chunks == 0:
|
||||
if not write_remainder:
|
||||
# Not enough to form a full chunk; keep buffered for next round
|
||||
# Re-buffer and return 0 written
|
||||
self._tables = [combined]
|
||||
return 0
|
||||
# Last flush: write the small remainder as a final file in worker_0
|
||||
worker_dir = os.path.join(self.out_dir, "worker_0")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
# Determine next index
|
||||
num_parquets = 0
|
||||
for _, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
chunk_path = os.path.join(worker_dir, f"data_chunk_{num_parquets}.parquet")
|
||||
temp_path = chunk_path + '.tmp'
|
||||
pq.write_table(combined, temp_path, compression=self.compression)
|
||||
if os.path.exists(chunk_path):
|
||||
os.remove(chunk_path)
|
||||
os.rename(temp_path, chunk_path)
|
||||
return num_samples
|
||||
|
||||
# Only write full chunks; keep remainder for next flush
|
||||
written_rows = total_chunks * self.samples_per_file
|
||||
remainder = num_samples - written_rows
|
||||
|
||||
table_to_write = combined.slice(0, written_rows)
|
||||
remainder_table = combined.slice(written_rows, remainder) if remainder > 0 else None
|
||||
if remainder_table is not None and len(remainder_table) > 0:
|
||||
if write_remainder:
|
||||
# Write the remainder as a final small file (worker_0)
|
||||
worker_dir = os.path.join(self.out_dir, "worker_0")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
num_parquets = 0
|
||||
for _, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
remainder_path = os.path.join(worker_dir,
|
||||
f"data_chunk_{num_parquets}.parquet")
|
||||
temp_path = remainder_path + '.tmp'
|
||||
pq.write_table(remainder_table,
|
||||
temp_path,
|
||||
compression=self.compression)
|
||||
if os.path.exists(remainder_path):
|
||||
os.remove(remainder_path)
|
||||
os.rename(temp_path, remainder_path)
|
||||
else:
|
||||
self._tables = [remainder_table]
|
||||
|
||||
# Parallel write by chunk ranges
|
||||
if num_workers is None:
|
||||
num_workers = min(multiprocessing.cpu_count(), max(total_chunks, 1))
|
||||
num_workers = max(int(num_workers), 1)
|
||||
chunks_per_worker = (total_chunks + num_workers - 1) // num_workers
|
||||
|
||||
work_ranges: list[tuple[int, int, pa.Table, int, str, int, str]] = []
|
||||
for worker_id in range(num_workers):
|
||||
start_chunk = worker_id * chunks_per_worker
|
||||
end_chunk = min((worker_id + 1) * chunks_per_worker, total_chunks)
|
||||
if start_chunk < end_chunk:
|
||||
work_ranges.append(
|
||||
(
|
||||
start_chunk,
|
||||
end_chunk,
|
||||
table_to_write,
|
||||
worker_id,
|
||||
self.out_dir,
|
||||
self.samples_per_file,
|
||||
self.compression,
|
||||
)
|
||||
)
|
||||
|
||||
written_total = 0
|
||||
if len(work_ranges) == 1:
|
||||
written_total += _process_chunk_range(work_ranges[0])
|
||||
return written_total
|
||||
|
||||
with ProcessPoolExecutor(max_workers=num_workers) as executor:
|
||||
futures = [executor.submit(_process_chunk_range, args) for args in work_ranges]
|
||||
for f in futures:
|
||||
written_total += f.result()
|
||||
return written_total + (len(remainder_table) if write_remainder and remainder_table is not None else 0)
|
||||
|
||||
|
||||
def _process_chunk_range(args: Any) -> int:
|
||||
"""Worker function to write a contiguous range of chunk files.
|
||||
|
||||
Args:
|
||||
args: Tuple containing
|
||||
- start_chunk (int): inclusive start chunk index
|
||||
- end_chunk (int): exclusive end chunk index
|
||||
- table (pa.Table): concatenated table containing all rows to write
|
||||
- worker_id (int): numeric worker identifier
|
||||
- output_dir (str): base output directory
|
||||
- samples_per_file (int): rows per chunk file
|
||||
- compression (str): compression codec for Parquet
|
||||
|
||||
Returns:
|
||||
int: Total number of rows written by this worker.
|
||||
"""
|
||||
start_chunk, end_chunk, table, worker_id, output_dir, samples_per_file, compression = args
|
||||
total_written = 0
|
||||
num_samples = len(table)
|
||||
|
||||
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
|
||||
# Offset to continue numbering if files exist
|
||||
num_parquets = 0
|
||||
for root, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
|
||||
for i in range(start_chunk, end_chunk):
|
||||
start_sample = i * samples_per_file
|
||||
end_sample = min((i + 1) * samples_per_file, num_samples)
|
||||
if end_sample <= start_sample:
|
||||
continue
|
||||
chunk = table.slice(start_sample, end_sample - start_sample)
|
||||
|
||||
chunk_path = os.path.join(worker_dir, f"data_chunk_{i + num_parquets}.parquet")
|
||||
temp_path = chunk_path + '.tmp'
|
||||
try:
|
||||
pq.write_table(chunk, temp_path, compression=compression)
|
||||
if os.path.exists(chunk_path):
|
||||
os.remove(chunk_path)
|
||||
os.rename(temp_path, chunk_path)
|
||||
total_written += len(chunk)
|
||||
except Exception:
|
||||
if os.path.exists(temp_path):
|
||||
os.remove(temp_path)
|
||||
raise
|
||||
|
||||
return total_written
|
||||
|
||||
|
||||
|
||||
+68
@@ -1,5 +1,7 @@
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
|
||||
|
||||
|
||||
@@ -120,3 +122,69 @@ def i2v_record_creator(batch: PreprocessBatch) -> list[dict[str, Any]]:
|
||||
})
|
||||
|
||||
return records
|
||||
|
||||
|
||||
def ode_text_only_record_creator(
|
||||
video_name: str, text_embedding: np.ndarray, caption: str,
|
||||
trajectory_latents: np.ndarray,
|
||||
trajectory_timesteps: np.ndarray) -> dict[str, Any]:
|
||||
"""Create a text-only ODE trajectory record matching pyarrow_schema_ode_trajectory_text_only.
|
||||
|
||||
Args:
|
||||
video_name: Base name/id for the sample (without extension).
|
||||
text_embedding: Text encoder output array [SeqLen, Dim].
|
||||
caption: Original text prompt.
|
||||
trajectory_latents: Collected trajectory latents array.
|
||||
trajectory_timesteps: Collected timesteps array.
|
||||
|
||||
Returns:
|
||||
dict suitable for records_to_table(…, pyarrow_schema_ode_trajectory_text_only)
|
||||
"""
|
||||
assert trajectory_latents is not None, "trajectory_latents is required"
|
||||
assert trajectory_timesteps is not None, "trajectory_timesteps is required"
|
||||
|
||||
record = {
|
||||
"id": f"text_{video_name}",
|
||||
"text_embedding_bytes": text_embedding.tobytes(),
|
||||
"text_embedding_shape": list(text_embedding.shape),
|
||||
"text_embedding_dtype": str(text_embedding.dtype),
|
||||
"file_name": video_name,
|
||||
"caption": caption,
|
||||
"media_type": "text",
|
||||
}
|
||||
|
||||
record.update({
|
||||
"trajectory_latents_bytes": trajectory_latents.tobytes(),
|
||||
"trajectory_latents_shape": list(trajectory_latents.shape),
|
||||
"trajectory_latents_dtype": str(trajectory_latents.dtype),
|
||||
})
|
||||
|
||||
record.update({
|
||||
"trajectory_timesteps_bytes": trajectory_timesteps.tobytes(),
|
||||
"trajectory_timesteps_shape": list(trajectory_timesteps.shape),
|
||||
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
|
||||
})
|
||||
|
||||
return record
|
||||
|
||||
|
||||
def text_only_record_creator(text_name: str, text_embedding: np.ndarray,
|
||||
caption: str) -> dict[str, Any]:
|
||||
"""Create a text-only record matching pyarrow_schema_text_only.
|
||||
|
||||
Args:
|
||||
text_name: Base id/name for the text sample.
|
||||
text_embedding: Text encoder output array [SeqLen, Dim].
|
||||
caption: Original text prompt.
|
||||
|
||||
Returns:
|
||||
dict suitable for records_to_table(…, pyarrow_schema_text_only)
|
||||
"""
|
||||
record = {
|
||||
"id": f"text_{text_name}",
|
||||
"text_embedding_bytes": text_embedding.tobytes(),
|
||||
"text_embedding_shape": list(text_embedding.shape),
|
||||
"text_embedding_dtype": str(text_embedding.dtype),
|
||||
"caption": caption,
|
||||
}
|
||||
return record
|
||||
@@ -50,6 +50,7 @@ pyarrow_schema_i2v = pa.schema([
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
|
||||
pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Image/Video VAE latents ---
|
||||
@@ -78,3 +79,40 @@ pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
|
||||
pyarrow_schema_ode_trajectory_text_only = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
# --- ODE Trajectory ---
|
||||
pa.field("trajectory_latents_bytes", pa.binary()),
|
||||
pa.field("trajectory_latents_shape", pa.list_(pa.int64())),
|
||||
pa.field("trajectory_latents_dtype", pa.string()),
|
||||
pa.field("trajectory_timesteps_bytes", pa.binary()),
|
||||
pa.field("trajectory_timesteps_shape", pa.list_(pa.int64())),
|
||||
pa.field("trajectory_timesteps_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # Always 'text' for text-only
|
||||
])
|
||||
|
||||
|
||||
pyarrow_schema_text_only = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("caption", pa.string()),
|
||||
])
|
||||
|
||||
@@ -628,3 +628,134 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
|
||||
"""Load state dict from checkpoint."""
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
|
||||
|
||||
class TextDataset(torch.utils.data.IterableDataset,
|
||||
torch.distributed.checkpoint.stateful.Stateful):
|
||||
"""
|
||||
Text-only dataset for processing prompts from a simple text file.
|
||||
|
||||
Assumes that data_merge_path is a text file with one prompt per line:
|
||||
A cat playing with a ball
|
||||
A dog running in the park
|
||||
A person cooking dinner
|
||||
...
|
||||
|
||||
This dataset processes text data through text encoding stages only.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
data_merge_path: str,
|
||||
args,
|
||||
start_idx: int = 0,
|
||||
seed: int = 42):
|
||||
self.data_merge_path = data_merge_path
|
||||
self.start_idx = start_idx
|
||||
self.args = args
|
||||
self.seed = seed
|
||||
|
||||
# Initialize tokenizer
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
|
||||
# Initialize text encoding stage
|
||||
self.text_encoding_stage = TextEncodingStage(
|
||||
tokenizer=tokenizer,
|
||||
text_max_length=args.text_max_length,
|
||||
cfg_rate=getattr(args, 'training_cfg_rate', 0.0),
|
||||
seed=self.seed)
|
||||
|
||||
# Process text data
|
||||
self.processed_batches = self._process_text_data()
|
||||
|
||||
def _load_text_data(self) -> list[str]:
|
||||
"""Load text prompts from file."""
|
||||
prompts = []
|
||||
with open(self.data_merge_path, 'r', encoding='utf-8') as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line: # Skip empty lines
|
||||
prompts.append(line)
|
||||
|
||||
logger.info(f"Loaded {len(prompts)} text prompts from {self.data_merge_path}")
|
||||
return prompts
|
||||
|
||||
def _process_text_data(self) -> list[PreprocessBatch]:
|
||||
"""Process the text prompts through text encoding stage."""
|
||||
raw_prompts = self._load_text_data()
|
||||
processed_batches = []
|
||||
|
||||
for idx, prompt in enumerate(raw_prompts):
|
||||
# Create a text-only batch with dummy path
|
||||
batch = PreprocessBatch(
|
||||
path=f"text_prompt_{idx}",
|
||||
cap=[prompt], # TextEncodingStage expects a list
|
||||
resolution=None,
|
||||
fps=None,
|
||||
duration=None,
|
||||
num_frames=0,
|
||||
sample_frame_index=None,
|
||||
sample_num_frames=0
|
||||
)
|
||||
|
||||
processed_batches.append(batch)
|
||||
|
||||
logger.info(f"Processed {len(processed_batches)} text batches")
|
||||
return processed_batches
|
||||
|
||||
def __iter__(self):
|
||||
"""Iterator for the dataset."""
|
||||
# Set up distributed sampling if needed
|
||||
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
||||
rank = torch.distributed.get_rank()
|
||||
world_size = torch.distributed.get_world_size()
|
||||
else:
|
||||
rank = 0
|
||||
world_size = 1
|
||||
|
||||
# Calculate chunk for this rank
|
||||
total_items = len(self.processed_batches)
|
||||
items_per_rank = math.ceil(total_items / world_size)
|
||||
start_idx = rank * items_per_rank + self.start_idx
|
||||
end_idx = min(start_idx + items_per_rank, total_items)
|
||||
|
||||
# Yield items for this rank
|
||||
for idx in range(start_idx, end_idx):
|
||||
if idx < len(self.processed_batches):
|
||||
yield self._get_item(idx)
|
||||
|
||||
def _get_item(self, idx: int) -> dict:
|
||||
"""Get a single processed text item."""
|
||||
batch = self.processed_batches[idx]
|
||||
|
||||
# Apply text encoding stage
|
||||
batch = self.text_encoding_stage.process(batch)
|
||||
|
||||
# Build result dictionary for text-only processing with required schema fields
|
||||
result = {
|
||||
"text": batch.text,
|
||||
"input_ids": batch.input_ids,
|
||||
"cond_mask": batch.cond_mask,
|
||||
"path": batch.path,
|
||||
# Required schema fields for ODE trajectory processing
|
||||
"id": f"text_{idx}",
|
||||
"file_name": batch.path,
|
||||
"caption": batch.text,
|
||||
"media_type": "text",
|
||||
"width": 1,
|
||||
"height": 1,
|
||||
"num_frames": 0,
|
||||
"duration_sec": 0.0,
|
||||
"fps": 0.0,
|
||||
}
|
||||
|
||||
return result
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
"""Return state dict for checkpointing."""
|
||||
return {"processed_batches": self.processed_batches}
|
||||
|
||||
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
|
||||
"""Load state dict from checkpoint."""
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
|
||||
@@ -5,7 +5,7 @@ import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
|
||||
def pad(t: torch.Tensor, padding_length: int) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Pad or crop an embedding [L, D] to exactly padding_length tokens.
|
||||
Return:
|
||||
|
||||
@@ -344,6 +344,9 @@ class VideoGenerator:
|
||||
"size": (target_height, target_width, batch.num_frames),
|
||||
"generation_time": gen_time,
|
||||
"logging_info": logging_info,
|
||||
"trajectory": output_batch.trajectory_latents,
|
||||
"trajectory_timesteps": output_batch.trajectory_timesteps,
|
||||
"trajectory_decoded": output_batch.trajectory_decoded,
|
||||
}
|
||||
|
||||
def set_lora_adapter(self,
|
||||
|
||||
@@ -175,6 +175,7 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
# - "SLIDING_TILE_ATTN" : use Sliding Tile Attention
|
||||
# - "VIDEO_SPARSE_ATTN": use Video Sparse Attention
|
||||
# - "SAGE_ATTN": use Sage Attention
|
||||
# - "SAGE_ATTN_THREE": use Sage Attention 3
|
||||
"FASTVIDEO_ATTENTION_BACKEND":
|
||||
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
|
||||
|
||||
|
||||
@@ -158,6 +158,7 @@ class FastVideoArgs:
|
||||
"transformer": True,
|
||||
"vae": True,
|
||||
})
|
||||
override_transformer_cls_name: str | None = None
|
||||
|
||||
# # DMD parameters
|
||||
# dmd_denoising_steps: List[int] | None = field(default=None)
|
||||
@@ -396,6 +397,12 @@ class FastVideoArgs:
|
||||
default=FastVideoArgs.enable_stage_verification,
|
||||
help="Enable input/output verification for pipeline stages",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--override-transformer-cls-name",
|
||||
type=str,
|
||||
default=FastVideoArgs.override_transformer_cls_name,
|
||||
help="Override transformer cls name",
|
||||
)
|
||||
# Add pipeline configuration arguments
|
||||
PipelineConfig.add_cli_args(parser)
|
||||
|
||||
@@ -605,6 +612,11 @@ class TrainingArgs(FastVideoArgs):
|
||||
pretrained_model_name_or_path: str = ""
|
||||
dit_model_name_or_path: str = ""
|
||||
|
||||
# DMD model paths - separate paths for each network
|
||||
generator_model_path: str = "" # path for generator (student) model
|
||||
real_score_model_path: str = "" # path for real score (teacher) model
|
||||
fake_score_model_path: str = "" # path for fake score (critic) model
|
||||
|
||||
# diffusion setting
|
||||
ema_decay: float = 0.0
|
||||
ema_start_step: int = 0
|
||||
@@ -627,6 +639,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
checkpoints_total_limit: int = 0
|
||||
checkpointing_steps: int = 0
|
||||
resume_from_checkpoint: str = "" # specify the checkpoint folder to resume from
|
||||
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
|
||||
|
||||
# optimizer & scheduler
|
||||
num_train_epochs: int = 0
|
||||
@@ -658,6 +671,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
linear_quadratic_threshold: float = 0.0
|
||||
linear_range: float = 0.0
|
||||
weight_decay: float = 0.0
|
||||
betas: str = "0.9,0.999" # betas for optimizer, format: "beta1,beta2"
|
||||
use_ema: bool = False
|
||||
multi_phased_distill_schedule: str = ""
|
||||
pred_decay_weight: float = 0.0
|
||||
@@ -678,16 +692,29 @@ class TrainingArgs(FastVideoArgs):
|
||||
|
||||
# distillation args
|
||||
generator_update_interval: int = 5
|
||||
dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic
|
||||
min_timestep_ratio: float = 0.2
|
||||
max_timestep_ratio: float = 0.98
|
||||
real_score_guidance_scale: float = 3.5
|
||||
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
|
||||
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
|
||||
fake_score_betas: str = "0.9,0.999" # betas for fake score optimizer, format: "beta1,beta2"
|
||||
training_state_checkpointing_steps: int = 0 # for resuming training
|
||||
weight_only_checkpointing_steps: int = 0 # for inference
|
||||
log_visualization: bool = False
|
||||
# simulate generator forward to match inference
|
||||
simulate_generator_forward: bool = False
|
||||
warp_denoising_step: bool = False
|
||||
|
||||
# Self-forcing specific arguments
|
||||
num_frame_per_block: int = 3
|
||||
independent_first_frame: bool = False
|
||||
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":
|
||||
@@ -789,6 +816,20 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=str,
|
||||
help="Directory to cache models")
|
||||
|
||||
# DMD model paths - separate paths for each network
|
||||
parser.add_argument(
|
||||
"--generator-model-path",
|
||||
type=str,
|
||||
help="Path to generator (student) model for DMD distillation")
|
||||
parser.add_argument(
|
||||
"--real-score-model-path",
|
||||
type=str,
|
||||
help="Path to real score (teacher) model for DMD distillation")
|
||||
parser.add_argument(
|
||||
"--fake-score-model-path",
|
||||
type=str,
|
||||
help="Path to fake score (critic) model for DMD distillation")
|
||||
|
||||
# Diffusion settings
|
||||
parser.add_argument("--ema-decay",
|
||||
type=float,
|
||||
@@ -859,6 +900,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--resume-from-checkpoint",
|
||||
type=str,
|
||||
help="Path to checkpoint to resume from")
|
||||
parser.add_argument(
|
||||
"--init-weights-from-safetensors",
|
||||
type=str,
|
||||
help="Path to safetensors file for initial weight loading")
|
||||
parser.add_argument("--logging-dir",
|
||||
type=str,
|
||||
help="Directory for logging")
|
||||
@@ -963,6 +1008,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
help="Linear quadratic threshold")
|
||||
parser.add_argument("--linear-range", type=float, help="Linear range")
|
||||
parser.add_argument("--weight-decay", type=float, help="Weight decay")
|
||||
parser.add_argument("--betas",
|
||||
type=str,
|
||||
default=TrainingArgs.betas,
|
||||
help="Betas for optimizer (format: 'beta1,beta2')")
|
||||
parser.add_argument("--use-ema",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use EMA")
|
||||
@@ -1013,6 +1062,13 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=int,
|
||||
default=TrainingArgs.generator_update_interval,
|
||||
help="Ratio of student updates to critic updates.")
|
||||
parser.add_argument(
|
||||
"--dfake-gen-update-ratio",
|
||||
type=int,
|
||||
default=TrainingArgs.dfake_gen_update_ratio,
|
||||
help=
|
||||
"Self-forcing: How often to train generator vs critic (train generator every N steps)."
|
||||
)
|
||||
parser.add_argument("--min-timestep-ratio",
|
||||
type=float,
|
||||
default=TrainingArgs.min_timestep_ratio,
|
||||
@@ -1029,6 +1085,11 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=float,
|
||||
default=TrainingArgs.fake_score_learning_rate,
|
||||
help="Learning rate for fake score transformer")
|
||||
parser.add_argument(
|
||||
"--fake-score-betas",
|
||||
type=str,
|
||||
default=TrainingArgs.fake_score_betas,
|
||||
help="Betas for fake score optimizer (format: 'beta1,beta2')")
|
||||
parser.add_argument(
|
||||
"--fake-score-lr-scheduler",
|
||||
type=str,
|
||||
@@ -1041,6 +1102,48 @@ class TrainingArgs(FastVideoArgs):
|
||||
"--simulate-generator-forward",
|
||||
action=StoreBoolean,
|
||||
help="Whether to simulate generator forward to match inference")
|
||||
parser.add_argument(
|
||||
"--warp-denoising-step",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Whether to warp denoising step according to the scheduler time shift"
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
parser.add_argument(
|
||||
"--num-frame-per-block",
|
||||
type=int,
|
||||
default=TrainingArgs.num_frame_per_block,
|
||||
help="Number of frames per block for causal generation")
|
||||
parser.add_argument(
|
||||
"--independent-first-frame",
|
||||
action=StoreBoolean,
|
||||
help="Whether the first frame is independent in causal generation")
|
||||
parser.add_argument(
|
||||
"--enable-gradient-masking",
|
||||
action=StoreBoolean,
|
||||
help="Whether to enable frame-level gradient masking")
|
||||
parser.add_argument(
|
||||
"--gradient-mask-last-n-frames",
|
||||
type=int,
|
||||
default=TrainingArgs.gradient_mask_last_n_frames,
|
||||
help="Number of last frames to enable gradients for")
|
||||
parser.add_argument(
|
||||
"--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
|
||||
|
||||
|
||||
@@ -77,9 +77,11 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
lora_A = self.lora_A.to_local()
|
||||
|
||||
if not self.merged and not self.disable_lora:
|
||||
delta = x @ (
|
||||
self.slice_lora_b_weights(lora_B.to(x, non_blocking=True))
|
||||
@ self.slice_lora_a_weights(lora_A.to(x, non_blocking=True)))
|
||||
lora_A_sliced = self.slice_lora_a_weights(
|
||||
lora_A.to(x, non_blocking=True))
|
||||
lora_B_sliced = self.slice_lora_b_weights(
|
||||
lora_B.to(x, non_blocking=True))
|
||||
delta = x @ lora_A_sliced.T @ lora_B_sliced.T
|
||||
if self.lora_alpha != self.lora_rank:
|
||||
delta = delta * (
|
||||
self.lora_alpha / self.lora_rank # type: ignore
|
||||
|
||||
@@ -147,6 +147,9 @@ 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"].detach()
|
||||
kv_cache["v"] = kv_cache["v"].detach()
|
||||
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
x = self.attn(
|
||||
@@ -375,7 +378,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
# Causal-specific
|
||||
self.block_mask = None
|
||||
self.num_frame_per_block = 1
|
||||
self.num_frame_per_block = config.arch_config.num_frames_per_block
|
||||
assert self.num_frame_per_block <= 3
|
||||
self.independent_first_frame = False
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
@@ -416,6 +416,11 @@ class TransformerLoader(ComponentLoader):
|
||||
"Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported.")
|
||||
|
||||
logger.info("transformer cls_name: %s", cls_name)
|
||||
if fastvideo_args.override_transformer_cls_name is not None:
|
||||
cls_name = fastvideo_args.override_transformer_cls_name
|
||||
logger.info("Overriding transformer cls_name to %s", cls_name)
|
||||
|
||||
fastvideo_args.model_paths["transformer"] = model_path
|
||||
|
||||
# Config from Diffusers supersedes fastvideo's model config
|
||||
@@ -430,8 +435,24 @@ class TransformerLoader(ComponentLoader):
|
||||
if not safetensors_list:
|
||||
raise ValueError(f"No safetensors files found in {model_path}")
|
||||
|
||||
logger.info("Loading model from %s safetensors files in %s",
|
||||
len(safetensors_list), model_path)
|
||||
# Check if we should use custom initialization weights
|
||||
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors', None)
|
||||
use_custom_weights = (custom_weights_path and os.path.exists(custom_weights_path) and
|
||||
fastvideo_args.training_mode and
|
||||
not hasattr(fastvideo_args, '_loading_teacher_critic_model'))
|
||||
|
||||
if use_custom_weights:
|
||||
logger.info("Using custom initialization weights from: %s", custom_weights_path)
|
||||
assert custom_weights_path is not None, "Custom initialization weights must be provided"
|
||||
if os.path.isdir(custom_weights_path):
|
||||
safetensors_list = glob.glob(
|
||||
os.path.join(str(custom_weights_path), "*.safetensors"))
|
||||
else:
|
||||
assert custom_weights_path.endswith(".safetensors"), "Custom initialization weights must be a safetensors file"
|
||||
safetensors_list = [custom_weights_path]
|
||||
|
||||
logger.info("Loading model from %s safetensors files: %s",
|
||||
len(safetensors_list), safetensors_list)
|
||||
|
||||
default_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.dit_precision]
|
||||
@@ -454,6 +475,7 @@ class TransformerLoader(ComponentLoader):
|
||||
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
|
||||
fsdp_inference=fastvideo_args.use_fsdp_inference,
|
||||
# TODO(will): make these configurable
|
||||
default_dtype=default_dtype,
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
output_dtype=None,
|
||||
@@ -463,9 +485,8 @@ class TransformerLoader(ComponentLoader):
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||
|
||||
dtypes = set(param.dtype for param in model.parameters())
|
||||
if len(dtypes) > 1:
|
||||
model = model.to(default_dtype)
|
||||
assert next(model.parameters()).dtype == default_dtype, "Model dtype does not match default dtype"
|
||||
|
||||
model = model.eval()
|
||||
return model
|
||||
|
||||
|
||||
@@ -62,6 +62,7 @@ def maybe_load_fsdp_model(
|
||||
device: torch.device,
|
||||
hsdp_replicate_dim: int,
|
||||
hsdp_shard_dim: int,
|
||||
default_dtype: torch.dtype,
|
||||
param_dtype: torch.dtype,
|
||||
reduce_dtype: torch.dtype,
|
||||
cpu_offload: bool = False,
|
||||
@@ -87,7 +88,8 @@ def maybe_load_fsdp_model(
|
||||
mp_policy=mp_policy,
|
||||
)
|
||||
|
||||
with set_default_dtype(param_dtype), torch.device("meta"):
|
||||
logger.info("Loading model with default_dtype: %s", default_dtype)
|
||||
with set_default_dtype(default_dtype), torch.device("meta"):
|
||||
model = model_cls(**init_params)
|
||||
|
||||
# Check if we should use FSDP
|
||||
@@ -125,7 +127,7 @@ def maybe_load_fsdp_model(
|
||||
model,
|
||||
weight_iterator,
|
||||
device,
|
||||
param_dtype,
|
||||
default_dtype,
|
||||
strict=True,
|
||||
cpu_offload=cpu_offload,
|
||||
param_names_mapping=param_names_mapping_fn,
|
||||
|
||||
@@ -635,8 +635,31 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
noise: torch.Tensor,
|
||||
timestep: torch.IntTensor,
|
||||
) -> torch.Tensor:
|
||||
|
||||
"""
|
||||
Args:
|
||||
clean_latent: the clean latent with shape [B, C, H, W],
|
||||
where B is batch_size or batch_size * num_frames
|
||||
noise: the noise with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
|
||||
|
||||
Returns:
|
||||
the corrupted latent with shape [B, C, H, W]
|
||||
"""
|
||||
# If timestep is [bs, num_frames]
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
assert timestep.numel() == clean_latent.shape[0]
|
||||
elif timestep.ndim == 1:
|
||||
# If timestep is [1]
|
||||
if timestep.shape[0] == 1:
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
else:
|
||||
assert timestep.numel() == clean_latent.shape[0]
|
||||
else:
|
||||
raise ValueError(f"[add_noise] Invalid timestep shape: {timestep.shape}")
|
||||
# timestep shape should be [B]
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
|
||||
@@ -22,8 +22,10 @@ class SelfForcingFlowMatchSchedulerOutput(BaseOutput):
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
|
||||
|
||||
config_name = "scheduler_config.json"
|
||||
order = 1
|
||||
@register_to_config
|
||||
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False, training=False):
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.shift = shift
|
||||
@@ -62,8 +64,15 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
def step(self, model_output: torch.FloatTensor, timestep: torch.FloatTensor, sample: torch.FloatTensor, to_final=False, return_dict=False, **kwargs):
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
elif timestep.ndim == 0:
|
||||
# handles the case where timestep is a scalar, this occurs when we
|
||||
# use this scheduler for ODE trajectory
|
||||
timestep = timestep.unsqueeze(0)
|
||||
|
||||
self.sigmas = self.sigmas.to(model_output.device)
|
||||
self.timesteps = self.timesteps.to(model_output.device)
|
||||
timestep = timestep.to(model_output.device)
|
||||
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
@@ -171,10 +171,10 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
|
||||
# timestep shape should be [B]
|
||||
dtype = pred_noise.dtype
|
||||
device = pred_noise.device
|
||||
pred_noise = pred_noise.float().to(device)
|
||||
noise_input_latent = noise_input_latent.float().to(device)
|
||||
sigmas = scheduler.sigmas.float().to(device)
|
||||
timesteps = scheduler.timesteps.float().to(device)
|
||||
pred_noise = pred_noise.double().to(device)
|
||||
noise_input_latent = noise_input_latent.double().to(device)
|
||||
sigmas = scheduler.sigmas.double().to(device)
|
||||
timesteps = scheduler.timesteps.double().to(device)
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
@@ -7,8 +7,6 @@ This module wires the causal DMD denoising stage into the modular pipeline.
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
|
||||
# isort: off
|
||||
@@ -28,10 +26,6 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
|
||||
@@ -40,7 +40,7 @@ class ComposedPipelineBase(ABC):
|
||||
_extra_config_module_map: dict[str, str] = {}
|
||||
training_args: TrainingArgs | None = None
|
||||
fastvideo_args: FastVideoArgs | TrainingArgs | None = None
|
||||
modules: dict[str, torch.nn.Module] = {}
|
||||
modules: dict[str, Any] = {}
|
||||
post_init_called: bool = False
|
||||
|
||||
# TODO(will): args should support both inference args and training args
|
||||
@@ -237,20 +237,19 @@ class ComposedPipelineBase(ABC):
|
||||
# remove keys that are not pipeline modules
|
||||
model_index.pop("_class_name")
|
||||
model_index.pop("_diffusers_version")
|
||||
# @TODO(Wei): Temporary hack
|
||||
if "boundary_ratio" in model_index and model_index[
|
||||
"boundary_ratio"] is not None:
|
||||
logger.info(
|
||||
"MoE pipeline detected. Adding transformer_2 to self.required_config_modules..."
|
||||
)
|
||||
self.required_config_modules.append("transformer_2")
|
||||
if fastvideo_args.boundary_ratio is None:
|
||||
logger.info(
|
||||
"MoE pipeline detected. Setting boundary ratio to %s",
|
||||
model_index["boundary_ratio"])
|
||||
fastvideo_args.boundary_ratio = model_index["boundary_ratio"]
|
||||
logger.info("MoE pipeline detected. Setting boundary ratio to %s",
|
||||
model_index["boundary_ratio"])
|
||||
fastvideo_args.pipeline_config.dit_config.boundary_ratio = model_index[
|
||||
"boundary_ratio"]
|
||||
|
||||
model_index.pop("boundary_ratio", None)
|
||||
# used by Wan2.2 ti2v
|
||||
model_index.pop("expand_timesteps", None)
|
||||
|
||||
# some sanity checks
|
||||
@@ -283,8 +282,8 @@ class ComposedPipelineBase(ABC):
|
||||
architecture) in model_index.items():
|
||||
if transformers_or_diffusers is None:
|
||||
logger.warning(
|
||||
"Module in model_index.json has null value, removing from required_config_modules"
|
||||
)
|
||||
"Module %s in model_index.json has null value, removing from required_config_modules",
|
||||
module_name)
|
||||
if module_name in self.required_config_modules:
|
||||
self.required_config_modules.remove(module_name)
|
||||
continue
|
||||
|
||||
@@ -129,6 +129,7 @@ class ForwardBatch:
|
||||
timesteps: torch.Tensor | None = None
|
||||
timestep: torch.Tensor | float | int | None = None
|
||||
step_index: int | None = None
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Scheduler parameters
|
||||
num_inference_steps: int = 50
|
||||
@@ -147,7 +148,12 @@ class ForwardBatch:
|
||||
modules: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# Final output (after pipeline completion)
|
||||
output: Any = None
|
||||
output: torch.Tensor | None = None
|
||||
return_trajectory_latents: bool = False
|
||||
return_trajectory_decoded: bool = False
|
||||
trajectory_timesteps: list[torch.Tensor] | None = None
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_decoded: list[torch.Tensor] | None = None
|
||||
|
||||
# Extra parameters that might be needed by specific pipeline implementations
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
@@ -206,6 +212,10 @@ class TrainingBatch:
|
||||
infos: list[dict[str, Any]] | None = None
|
||||
mask_lat_size: torch.Tensor | None = None
|
||||
|
||||
# ODE trajectory supervision
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_timesteps: torch.Tensor | None = None
|
||||
|
||||
# Transformer inputs
|
||||
noisy_model_input: torch.Tensor | None = None
|
||||
timesteps: torch.Tensor | None = None
|
||||
@@ -236,6 +246,7 @@ class TrainingBatch:
|
||||
fake_score_loss: float = 0.0
|
||||
|
||||
dmd_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import multiprocessing
|
||||
import os
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
@@ -12,6 +10,8 @@ from torch.utils.data import DataLoader
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset import getdataset
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.dataset.preprocessing_datasets import PreprocessBatch
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -54,10 +54,14 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"""Get additional features specific to the pipeline type. Override in subclasses."""
|
||||
return {}
|
||||
|
||||
def get_schema_fields(self) -> list[str]:
|
||||
"""Get the schema fields for the pipeline type. Override in subclasses."""
|
||||
def get_pyarrow_schema(self) -> pa.Schema:
|
||||
"""Return the PyArrow schema for this pipeline. Must be overridden."""
|
||||
raise NotImplementedError
|
||||
|
||||
def get_schema_fields(self) -> list[str]:
|
||||
"""Get the schema fields for the pipeline type."""
|
||||
return [f.name for f in self.get_pyarrow_schema()]
|
||||
|
||||
def create_record_for_schema(self,
|
||||
preprocess_batch: PreprocessBatch,
|
||||
schema: pa.Schema,
|
||||
@@ -400,166 +404,22 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
batch_data.append(record)
|
||||
|
||||
if batch_data:
|
||||
# Add progress bar for writing to Parquet dataset
|
||||
write_pbar = tqdm(total=1,
|
||||
desc="Writing to Parquet dataset",
|
||||
unit="batch")
|
||||
# Convert batch data to PyArrow arrays
|
||||
arrays = []
|
||||
for field in self.get_schema_fields():
|
||||
if field.endswith('_bytes'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.binary()))
|
||||
elif field.endswith('_shape'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.list_(pa.int32())))
|
||||
elif field in ['width', 'height', 'num_frames']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.int32()))
|
||||
elif field in ['duration_sec', 'fps']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.float32()))
|
||||
else:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data]))
|
||||
|
||||
table = pa.Table.from_arrays(arrays,
|
||||
names=self.get_schema_fields())
|
||||
table = records_to_table(batch_data, self.get_pyarrow_schema())
|
||||
write_pbar.update(1)
|
||||
write_pbar.close()
|
||||
|
||||
# Store the table in a list for later processing
|
||||
if not hasattr(self, 'all_tables'):
|
||||
self.all_tables = []
|
||||
self.all_tables.append(table)
|
||||
|
||||
if not hasattr(self, 'dataset_writer'):
|
||||
self.dataset_writer = ParquetDatasetWriter(
|
||||
out_dir=combined_parquet_dir,
|
||||
samples_per_file=args.samples_per_file,
|
||||
)
|
||||
self.dataset_writer.append_table(table)
|
||||
logger.info("Collected batch with %s samples", len(table))
|
||||
|
||||
if num_processed_samples >= args.flush_frequency:
|
||||
self._flush_tables(num_processed_samples, args,
|
||||
combined_parquet_dir)
|
||||
written = self.dataset_writer.flush()
|
||||
logger.info("Flushed %s samples to parquet", written)
|
||||
num_processed_samples = 0
|
||||
self.all_tables = []
|
||||
|
||||
def _flush_tables(self, num_processed_samples: int, args,
|
||||
combined_parquet_dir: str):
|
||||
"""Flush collected tables to disk."""
|
||||
assert hasattr(self, 'all_tables') and self.all_tables
|
||||
print(f"Combining {len(self.all_tables)} batches...")
|
||||
combined_table = pa.concat_tables(self.all_tables)
|
||||
assert len(combined_table) == num_processed_samples
|
||||
print(f"Total samples collected: {len(combined_table)}")
|
||||
|
||||
# Calculate total number of chunks needed, discarding remainder
|
||||
total_chunks = max(num_processed_samples // args.samples_per_file, 1)
|
||||
|
||||
print(f"Fixed samples per parquet file: {args.samples_per_file}")
|
||||
print(f"Total number of parquet files: {total_chunks}")
|
||||
print(
|
||||
f"Total samples to be processed: {total_chunks * args.samples_per_file} (discarding {num_processed_samples % args.samples_per_file} samples)"
|
||||
)
|
||||
|
||||
# Split work among processes
|
||||
num_workers = int(min(multiprocessing.cpu_count(), total_chunks))
|
||||
chunks_per_worker = (total_chunks + num_workers - 1) // num_workers
|
||||
|
||||
print(f"Using {num_workers} workers to process {total_chunks} chunks")
|
||||
logger.info("Chunks per worker: %s", chunks_per_worker)
|
||||
|
||||
# Prepare work ranges
|
||||
work_ranges = []
|
||||
for i in range(num_workers):
|
||||
start_idx = i * chunks_per_worker
|
||||
end_idx = min((i + 1) * chunks_per_worker, total_chunks)
|
||||
if start_idx < total_chunks:
|
||||
work_ranges.append(
|
||||
(start_idx, end_idx, combined_table, i,
|
||||
combined_parquet_dir, args.samples_per_file))
|
||||
|
||||
total_written = 0
|
||||
failed_ranges = []
|
||||
with ProcessPoolExecutor(max_workers=num_workers) as executor:
|
||||
futures = {
|
||||
executor.submit(self.process_chunk_range, work_range):
|
||||
work_range
|
||||
for work_range in work_ranges
|
||||
}
|
||||
for future in tqdm(futures, desc="Processing chunks"):
|
||||
try:
|
||||
written = future.result()
|
||||
total_written += written
|
||||
logger.info("Processed chunk with %s samples", written)
|
||||
except Exception as e:
|
||||
work_range = futures[future]
|
||||
failed_ranges.append(work_range)
|
||||
logger.error("Failed to process range %s-%s: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
# Retry failed ranges sequentially
|
||||
if failed_ranges:
|
||||
logger.warning("Retrying %s failed ranges sequentially",
|
||||
len(failed_ranges))
|
||||
for work_range in failed_ranges:
|
||||
try:
|
||||
total_written += self.process_chunk_range(work_range)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to process range %s-%s after retry: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
logger.info("Total samples written: %s", total_written)
|
||||
|
||||
@staticmethod
|
||||
def process_chunk_range(args: Any) -> int:
|
||||
start_idx, end_idx, table, worker_id, output_dir, samples_per_file = args
|
||||
try:
|
||||
total_written = 0
|
||||
num_samples = len(table)
|
||||
|
||||
# Create worker-specific subdirectory
|
||||
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
|
||||
# Check how many files there are already in the dir, and update i accordingly
|
||||
num_parquets = 0
|
||||
for root, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
|
||||
for i in range(start_idx, end_idx):
|
||||
start_sample = i * samples_per_file
|
||||
end_sample = min((i + 1) * samples_per_file, num_samples)
|
||||
chunk = table.slice(start_sample, end_sample - start_sample)
|
||||
|
||||
# Create chunk file in worker's directory
|
||||
chunk_path = os.path.join(
|
||||
worker_dir, f"data_chunk_{i + num_parquets}.parquet")
|
||||
temp_path = chunk_path + '.tmp'
|
||||
|
||||
try:
|
||||
# Write to temporary file
|
||||
pq.write_table(chunk, temp_path, compression='zstd')
|
||||
|
||||
# Rename temporary file to final file
|
||||
if os.path.exists(chunk_path):
|
||||
os.remove(
|
||||
chunk_path) # Remove existing file if it exists
|
||||
os.rename(temp_path, chunk_path)
|
||||
|
||||
total_written += len(chunk)
|
||||
except Exception as e:
|
||||
# Clean up temporary file if it exists
|
||||
if os.path.exists(temp_path):
|
||||
os.remove(temp_path)
|
||||
raise e
|
||||
|
||||
return total_written
|
||||
except Exception as e:
|
||||
logger.error("Error processing chunks %s-%s for worker %s: %s",
|
||||
start_idx, end_idx, worker_id, str(e))
|
||||
raise
|
||||
|
||||
@@ -40,9 +40,9 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
image_processor=self.get_module("image_processor"),
|
||||
))
|
||||
|
||||
def get_schema_fields(self) -> list[str]:
|
||||
"""Get the schema fields for I2V pipeline."""
|
||||
return [f.name for f in pyarrow_schema_i2v]
|
||||
def get_pyarrow_schema(self):
|
||||
"""Return the PyArrow schema for I2V pipeline."""
|
||||
return pyarrow_schema_i2v
|
||||
|
||||
def get_extra_features(self, valid_data: dict[str, Any],
|
||||
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
|
||||
|
||||
@@ -0,0 +1,323 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
ODE Trajectory Data Preprocessing pipeline implementation.
|
||||
|
||||
This module contains an implementation of the ODE Trajectory Data Preprocessing pipeline
|
||||
using the modular pipeline architecture.
|
||||
|
||||
Sec 4.3 of CausVid paper: https://arxiv.org/pdf/2412.07772
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
|
||||
import pyarrow as pa
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset import gettextdataset
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.dataset.dataloader.record_schema import (
|
||||
ode_text_only_record_creator)
|
||||
from fastvideo.dataset.dataloader.schema import (
|
||||
pyarrow_schema_ode_trajectory_text_only)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
|
||||
SelfForcingFlowMatchScheduler)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
"""ODE Trajectory preprocessing pipeline implementation."""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
preprocess_dataloader: StatefulDataLoader
|
||||
preprocess_loader_iter: Iterator[dict[str, Any]]
|
||||
pbar: Any
|
||||
num_processed_samples: int
|
||||
|
||||
def get_pyarrow_schema(self) -> pa.Schema:
|
||||
"""Return the PyArrow schema for ODE Trajectory pipeline."""
|
||||
return pyarrow_schema_ode_trajectory_text_only
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
assert fastvideo_args.pipeline_config.flow_shift == 5
|
||||
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
sigma_min=0.0,
|
||||
extra_one_step=True)
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=48,
|
||||
denoising_strength=1.0)
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer", None)))
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
pipeline=self,
|
||||
))
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def preprocess_text_and_trajectory(self, fastvideo_args: FastVideoArgs,
|
||||
args):
|
||||
"""Preprocess text-only data and generate trajectory information."""
|
||||
|
||||
for batch_idx, data in enumerate(self.pbar):
|
||||
if data is None:
|
||||
continue
|
||||
|
||||
with torch.inference_mode():
|
||||
# For text-only processing, we only need text data
|
||||
# Filter out samples without text
|
||||
valid_indices = []
|
||||
for i, text in enumerate(data["text"]):
|
||||
if text and text.strip(): # Check if text is not empty
|
||||
valid_indices.append(i)
|
||||
self.num_processed_samples += len(valid_indices)
|
||||
|
||||
if not valid_indices:
|
||||
continue
|
||||
|
||||
# Create new batch with only valid samples (text-only)
|
||||
valid_data = {
|
||||
"text": [data["text"][i] for i in valid_indices],
|
||||
"path": [data["path"][i] for i in valid_indices],
|
||||
}
|
||||
|
||||
# Add fps and duration if available in data
|
||||
if "fps" in data:
|
||||
valid_data["fps"] = [data["fps"][i] for i in valid_indices]
|
||||
if "duration" in data:
|
||||
valid_data["duration"] = [
|
||||
data["duration"][i] for i in valid_indices
|
||||
]
|
||||
|
||||
batch_captions = valid_data["text"]
|
||||
# Encode text using the standalone TextEncodingStage API
|
||||
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
batch_captions,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
prompt_embeds = prompt_embeds_list[0]
|
||||
prompt_attention_masks = prompt_masks_list[0]
|
||||
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
|
||||
|
||||
sampling_params = SamplingParam.from_pretrained(args.model_path)
|
||||
|
||||
# encode negative prompt for trajectory collection
|
||||
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
|
||||
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
sampling_params.negative_prompt,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
negative_prompt_embed = negative_prompt_embeds_list[0][0]
|
||||
negative_prompt_attention_mask = negative_prompt_masks_list[
|
||||
0][0]
|
||||
else:
|
||||
negative_prompt_embed = None
|
||||
negative_prompt_attention_mask = None
|
||||
|
||||
trajectory_latents = []
|
||||
trajectory_timesteps = []
|
||||
trajectory_decoded = []
|
||||
|
||||
for i, (prompt_embed, prompt_attention_mask) in enumerate(
|
||||
zip(prompt_embeds, prompt_attention_masks,
|
||||
strict=False)):
|
||||
prompt_embed = prompt_embed.unsqueeze(0)
|
||||
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
|
||||
|
||||
# Collect the trajectory data (text-to-video generation)
|
||||
batch = ForwardBatch(**shallow_asdict(sampling_params), )
|
||||
batch.prompt_embeds = [prompt_embed]
|
||||
batch.prompt_attention_mask = [prompt_attention_mask]
|
||||
batch.negative_prompt_embeds = [negative_prompt_embed]
|
||||
batch.negative_attention_mask = [
|
||||
negative_prompt_attention_mask
|
||||
]
|
||||
batch.num_inference_steps = 48
|
||||
batch.return_trajectory_latents = True
|
||||
# Enabling this will save the decoded trajectory videos.
|
||||
# Used for debugging.
|
||||
batch.return_trajectory_decoded = False
|
||||
batch.height = args.max_height
|
||||
batch.width = args.max_width
|
||||
batch.fps = args.train_fps
|
||||
batch.guidance_scale = 6.0
|
||||
batch.do_classifier_free_guidance = True
|
||||
|
||||
result_batch = self.input_validation_stage(
|
||||
batch, fastvideo_args)
|
||||
result_batch = self.timestep_preparation_stage(
|
||||
batch, fastvideo_args)
|
||||
result_batch = self.latent_preparation_stage(
|
||||
result_batch, fastvideo_args)
|
||||
result_batch = self.denoising_stage(result_batch,
|
||||
fastvideo_args)
|
||||
result_batch = self.decoding_stage(result_batch,
|
||||
fastvideo_args)
|
||||
|
||||
trajectory_latents.append(
|
||||
result_batch.trajectory_latents.cpu())
|
||||
trajectory_timesteps.append(
|
||||
result_batch.trajectory_timesteps.cpu())
|
||||
trajectory_decoded.append(result_batch.trajectory_decoded)
|
||||
|
||||
# Prepare extra features for text-only processing
|
||||
extra_features = {
|
||||
"trajectory_latents": trajectory_latents,
|
||||
"trajectory_timesteps": trajectory_timesteps
|
||||
}
|
||||
|
||||
if batch.return_trajectory_decoded:
|
||||
for i, decoded_frames in enumerate(trajectory_decoded):
|
||||
for j, decoded_frame in enumerate(decoded_frames):
|
||||
save_decoded_latents_as_video(
|
||||
decoded_frame,
|
||||
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
|
||||
args.train_fps)
|
||||
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data: list[dict[str, Any]] = []
|
||||
|
||||
# Add progress bar for saving outputs
|
||||
save_pbar = tqdm(enumerate(valid_data["path"]),
|
||||
desc="Saving outputs",
|
||||
unit="item",
|
||||
leave=False)
|
||||
|
||||
for idx, video_path in save_pbar:
|
||||
video_name = os.path.basename(video_path).split(".")[0]
|
||||
|
||||
# Convert tensors to numpy arrays
|
||||
text_embedding = prompt_embeds[idx].cpu().numpy()
|
||||
|
||||
# Get extra features for this sample
|
||||
sample_extra_features = {}
|
||||
if extra_features:
|
||||
for key, value in extra_features.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
sample_extra_features[key] = value[idx].cpu(
|
||||
).numpy()
|
||||
else:
|
||||
assert isinstance(value, list)
|
||||
if isinstance(value[idx], torch.Tensor):
|
||||
sample_extra_features[key] = value[idx].cpu(
|
||||
).float().numpy()
|
||||
else:
|
||||
sample_extra_features[key] = value[idx]
|
||||
|
||||
# Create record for Parquet dataset (text-only ODE schema)
|
||||
record: dict[str, Any] = ode_text_only_record_creator(
|
||||
video_name=video_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=valid_data["text"][idx],
|
||||
trajectory_latents=sample_extra_features[
|
||||
"trajectory_latents"],
|
||||
trajectory_timesteps=sample_extra_features[
|
||||
"trajectory_timesteps"],
|
||||
)
|
||||
batch_data.append(record)
|
||||
|
||||
if batch_data:
|
||||
write_pbar = tqdm(total=1,
|
||||
desc="Writing to Parquet dataset",
|
||||
unit="batch")
|
||||
table = records_to_table(batch_data,
|
||||
self.get_pyarrow_schema())
|
||||
write_pbar.update(1)
|
||||
write_pbar.close()
|
||||
|
||||
if not hasattr(self, 'dataset_writer'):
|
||||
self.dataset_writer = ParquetDatasetWriter(
|
||||
out_dir=self.combined_parquet_dir,
|
||||
samples_per_file=args.samples_per_file,
|
||||
)
|
||||
self.dataset_writer.append_table(table)
|
||||
|
||||
logger.info("Collected batch with %s samples", len(table))
|
||||
|
||||
if self.num_processed_samples >= args.flush_frequency:
|
||||
written = self.dataset_writer.flush()
|
||||
logger.info("Flushed %s samples to parquet", written)
|
||||
self.num_processed_samples = 0
|
||||
|
||||
# Final flush for any remaining samples
|
||||
if hasattr(self, 'dataset_writer'):
|
||||
written = self.dataset_writer.flush(write_remainder=True)
|
||||
if written:
|
||||
logger.info("Final flush wrote %s samples", written)
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
|
||||
self.local_rank = int(os.getenv("RANK", 0))
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
# Create directory for combined data
|
||||
self.combined_parquet_dir = os.path.join(args.output_dir,
|
||||
"combined_parquet_dataset")
|
||||
os.makedirs(self.combined_parquet_dir, exist_ok=True)
|
||||
|
||||
# Loading dataset
|
||||
train_dataset = gettextdataset(args)
|
||||
|
||||
self.preprocess_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
batch_size=args.preprocess_video_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
|
||||
|
||||
self.num_processed_samples = 0
|
||||
# Add progress bar for video preprocessing
|
||||
self.pbar = tqdm(self.preprocess_loader_iter,
|
||||
desc="Processing videos",
|
||||
unit="batch",
|
||||
disable=self.local_rank != 0)
|
||||
|
||||
# Initialize class variables for data sharing
|
||||
self.video_data: dict[str, Any] = {} # Store video metadata and paths
|
||||
self.latent_data: dict[str, Any] = {} # Store latent tensors
|
||||
self.preprocess_text_and_trajectory(fastvideo_args, args)
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_ODE_Trajectory
|
||||
@@ -15,9 +15,9 @@ class PreprocessPipeline_T2V(BasePreprocessPipeline):
|
||||
|
||||
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
|
||||
|
||||
def get_schema_fields(self):
|
||||
"""Get the schema fields for T2V pipeline."""
|
||||
return [f.name for f in pyarrow_schema_t2v]
|
||||
def get_pyarrow_schema(self):
|
||||
"""Return the PyArrow schema for T2V pipeline."""
|
||||
return pyarrow_schema_t2v
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_T2V
|
||||
|
||||
@@ -0,0 +1,184 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Text-only Data Preprocessing pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Text-only Data Preprocessing pipeline
|
||||
using the modular pipeline architecture, based on the ODE Trajectory preprocessing.
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset import gettextdataset
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.dataset.dataloader.record_schema import text_only_record_creator
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_text_only
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.pipelines.stages import TextEncodingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PreprocessPipeline_Text(BasePreprocessPipeline):
|
||||
"""Text-only preprocessing pipeline implementation."""
|
||||
|
||||
_required_config_modules = ["text_encoder", "tokenizer"]
|
||||
preprocess_dataloader: StatefulDataLoader
|
||||
preprocess_loader_iter: Iterator[dict[str, Any]]
|
||||
pbar: Any
|
||||
num_processed_samples: int = 0
|
||||
|
||||
def get_pyarrow_schema(self):
|
||||
"""Return the PyArrow schema for text-only pipeline."""
|
||||
return pyarrow_schema_text_only
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
def preprocess_text_only(self, fastvideo_args: FastVideoArgs, args):
|
||||
"""Preprocess text-only data."""
|
||||
|
||||
for batch_idx, data in enumerate(self.pbar):
|
||||
if data is None:
|
||||
continue
|
||||
|
||||
with torch.inference_mode():
|
||||
# For text-only processing, we only need text data
|
||||
# Filter out samples without text
|
||||
valid_indices = []
|
||||
for i, text in enumerate(data["text"]):
|
||||
if text and text.strip(): # Check if text is not empty
|
||||
valid_indices.append(i)
|
||||
self.num_processed_samples += len(valid_indices)
|
||||
|
||||
if not valid_indices:
|
||||
continue
|
||||
|
||||
# Create new batch with only valid samples (text-only)
|
||||
valid_data = {
|
||||
"text": [data["text"][i] for i in valid_indices],
|
||||
"path": [data["path"][i] for i in valid_indices],
|
||||
}
|
||||
|
||||
batch_captions = valid_data["text"]
|
||||
# Encode text using the standalone TextEncodingStage API
|
||||
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
batch_captions,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
prompt_embeds = prompt_embeds_list[0]
|
||||
prompt_attention_masks = prompt_masks_list[0]
|
||||
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
|
||||
|
||||
logger.info("===== prompt_embeds: %s", prompt_embeds.shape)
|
||||
logger.info("===== prompt_attention_masks: %s",
|
||||
prompt_attention_masks.shape)
|
||||
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data = []
|
||||
|
||||
# Add progress bar for saving outputs
|
||||
save_pbar = tqdm(enumerate(valid_data["path"]),
|
||||
desc="Saving outputs",
|
||||
unit="item",
|
||||
leave=False)
|
||||
|
||||
for idx, text_path in save_pbar:
|
||||
text_name = os.path.basename(text_path).split(".")[0]
|
||||
|
||||
# Convert tensors to numpy arrays
|
||||
text_embedding = prompt_embeds[idx].cpu().numpy()
|
||||
|
||||
# Create record for Parquet dataset (text-only schema)
|
||||
record = text_only_record_creator(
|
||||
text_name=text_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=valid_data["text"][idx],
|
||||
)
|
||||
batch_data.append(record)
|
||||
|
||||
if batch_data:
|
||||
write_pbar = tqdm(total=1,
|
||||
desc="Writing to Parquet dataset",
|
||||
unit="batch")
|
||||
table = records_to_table(batch_data,
|
||||
pyarrow_schema_text_only)
|
||||
write_pbar.update(1)
|
||||
write_pbar.close()
|
||||
|
||||
if not hasattr(self, 'dataset_writer'):
|
||||
self.dataset_writer = ParquetDatasetWriter(
|
||||
out_dir=self.combined_parquet_dir,
|
||||
samples_per_file=args.samples_per_file,
|
||||
)
|
||||
self.dataset_writer.append_table(table)
|
||||
|
||||
logger.info("Collected batch with %s samples", len(table))
|
||||
|
||||
if self.num_processed_samples >= args.flush_frequency:
|
||||
written = self.dataset_writer.flush()
|
||||
logger.info("Flushed %s samples to parquet", written)
|
||||
self.num_processed_samples = 0
|
||||
|
||||
# Final flush for any remaining samples
|
||||
if hasattr(self, 'dataset_writer'):
|
||||
written = self.dataset_writer.flush(write_remainder=True)
|
||||
if written:
|
||||
logger.info("Final flush wrote %s samples", written)
|
||||
|
||||
# Text-only record creation moved to fastvideo.dataset.dataloader.record_schema
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
|
||||
self.local_rank = int(os.getenv("RANK", 0))
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
# Create directory for combined data
|
||||
self.combined_parquet_dir = os.path.join(args.output_dir,
|
||||
"combined_parquet_dataset")
|
||||
os.makedirs(self.combined_parquet_dir, exist_ok=True)
|
||||
|
||||
# Loading text dataset
|
||||
train_dataset = gettextdataset(args)
|
||||
|
||||
self.preprocess_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
batch_size=args.preprocess_video_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
|
||||
|
||||
self.num_processed_samples = 0
|
||||
# Add progress bar for text preprocessing
|
||||
self.pbar = tqdm(self.preprocess_loader_iter,
|
||||
desc="Processing text",
|
||||
unit="batch",
|
||||
disable=self.local_rank != 0)
|
||||
|
||||
# Initialize class variables for data sharing
|
||||
self.text_data: dict[str, Any] = {} # Store text metadata and paths
|
||||
|
||||
self.preprocess_text_only(fastvideo_args, args)
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_Text
|
||||
@@ -1,5 +1,6 @@
|
||||
import argparse
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
@@ -9,8 +10,12 @@ from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v import (
|
||||
PreprocessPipeline_I2V)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (
|
||||
PreprocessPipeline_ODE_Trajectory)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
|
||||
PreprocessPipeline_T2V)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
|
||||
PreprocessPipeline_Text)
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -21,12 +26,22 @@ def main(args) -> None:
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
num_gpus = int(os.environ["WORLD_SIZE"])
|
||||
assert num_gpus == 1, "Only support 1 GPU"
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
|
||||
kwargs = {
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
|
||||
}
|
||||
|
||||
kwargs: dict[str, Any] = {}
|
||||
if args.preprocess_task == "text_only":
|
||||
kwargs = {
|
||||
"text_encoder_cpu_offload": False,
|
||||
}
|
||||
else:
|
||||
# Full config for video/image processing
|
||||
kwargs = {
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
|
||||
}
|
||||
pipeline_config.update_config_from_dict(kwargs)
|
||||
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=args.model_path,
|
||||
num_gpus=get_world_size(),
|
||||
@@ -35,7 +50,23 @@ def main(args) -> None:
|
||||
text_encoder_cpu_offload=False,
|
||||
pipeline_config=pipeline_config,
|
||||
)
|
||||
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
|
||||
if args.preprocess_task == "t2v":
|
||||
PreprocessPipeline = PreprocessPipeline_T2V
|
||||
elif args.preprocess_task == "i2v":
|
||||
PreprocessPipeline = PreprocessPipeline_I2V
|
||||
elif args.preprocess_task == "text_only":
|
||||
PreprocessPipeline = PreprocessPipeline_Text
|
||||
elif args.preprocess_task == "ode_trajectory":
|
||||
assert args.flow_shift is not None, "flow_shift is required for ode_trajectory"
|
||||
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
|
||||
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
|
||||
else:
|
||||
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
|
||||
f"Valid options: t2v, i2v, ode_trajectory, text_only")
|
||||
|
||||
logger.info("Preprocess task: %s using %s", args.preprocess_task,
|
||||
PreprocessPipeline.__name__)
|
||||
|
||||
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
|
||||
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
|
||||
|
||||
@@ -74,7 +105,12 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--preprocess_task", type=str, default="t2v")
|
||||
parser.add_argument("--flow_shift", type=float, default=None)
|
||||
parser.add_argument("--preprocess_task",
|
||||
type=str,
|
||||
default="t2v",
|
||||
choices=["t2v", "i2v", "text_only", "ode_trajectory"],
|
||||
help="Type of preprocessing task to run")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
|
||||
@@ -50,6 +50,63 @@ class DecodingStage(PipelineStage):
|
||||
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latents: torch.Tensor,
|
||||
fastvideo_args: FastVideoArgs) -> torch.Tensor:
|
||||
"""
|
||||
Decode latent representations into pixel space using VAE.
|
||||
|
||||
Args:
|
||||
latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents)
|
||||
fastvideo_args: Configuration containing:
|
||||
- disable_autocast: Whether to disable automatic mixed precision (default: False)
|
||||
- pipeline_config.vae_precision: VAE computation precision ("fp32", "fp16", "bf16")
|
||||
- pipeline_config.vae_tiling: Whether to enable VAE tiling for memory efficiency
|
||||
|
||||
Returns:
|
||||
Decoded video tensor with shape (batch, channels, frames, height, width),
|
||||
normalized to [0, 1] range and moved to CPU as float32
|
||||
"""
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
latents = latents.to(get_local_torch_device())
|
||||
|
||||
# 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
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# 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):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
# Decode latents
|
||||
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:
|
||||
latents = latents.to(vae_dtype)
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
return image
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
@@ -59,13 +116,28 @@ class DecodingStage(PipelineStage):
|
||||
"""
|
||||
Decode latent representations into pixel space.
|
||||
|
||||
This method processes the batch through the VAE decoder, converting latent
|
||||
representations to pixel-space video/images. It also optionally decodes
|
||||
trajectory latents for visualization purposes.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
batch: The current batch containing:
|
||||
- latents: Tensor to decode (batch, channels, frames, height_latents, width_latents)
|
||||
- return_trajectory_decoded (optional): Flag to decode trajectory latents
|
||||
- trajectory_latents (optional): Latents at different timesteps
|
||||
- trajectory_timesteps (optional): Corresponding timesteps
|
||||
fastvideo_args: Configuration containing:
|
||||
- output_type: "latent" to skip decoding, otherwise decode to pixels
|
||||
- vae_cpu_offload: Whether to offload VAE to CPU after decoding
|
||||
- model_loaded: Track VAE loading state
|
||||
- model_paths: Path to VAE model if loading needed
|
||||
|
||||
Returns:
|
||||
The batch with decoded outputs.
|
||||
Modified batch with:
|
||||
- output: Decoded frames (batch, channels, frames, height, width) as CPU float32
|
||||
- trajectory_decoded (if requested): List of decoded frames per timestep
|
||||
"""
|
||||
# load vae if not already loaded (used for memory constrained devices)
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
if not fastvideo_args.model_loaded["vae"]:
|
||||
loader = VAELoader()
|
||||
@@ -75,58 +147,29 @@ class DecodingStage(PipelineStage):
|
||||
pipeline.add_module("vae", self.vae)
|
||||
fastvideo_args.model_loaded["vae"] = True
|
||||
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
latents = batch.latents
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
if latents is None:
|
||||
raise ValueError("Latents must be provided")
|
||||
|
||||
# Skip decoding if output type is latent
|
||||
if fastvideo_args.output_type == "latent":
|
||||
image = latents
|
||||
frames = batch.latents
|
||||
else:
|
||||
# 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
|
||||
frames = self.decode(batch.latents, fastvideo_args)
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# 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):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
# Decode latents
|
||||
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:
|
||||
latents = latents.to(vae_dtype)
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
# decode trajectory latents if needed
|
||||
if batch.return_trajectory_decoded:
|
||||
batch.trajectory_decoded = []
|
||||
assert batch.trajectory_latents is not None, "batch should have trajectory latents"
|
||||
for idx in range(batch.trajectory_latents.shape[1]):
|
||||
# batch.trajectory_latents is [batch_size, timesteps, channels, frames, height, width]
|
||||
cur_latent = batch.trajectory_latents[:, idx, :, :, :, :]
|
||||
cur_timestep = batch.trajectory_timesteps[idx]
|
||||
logger.info("decoding trajectory latent for timestep: %s",
|
||||
cur_timestep)
|
||||
decoded_frames = self.decode(cur_latent, fastvideo_args)
|
||||
batch.trajectory_decoded.append(decoded_frames.cpu().float())
|
||||
|
||||
# Convert to CPU float32 for compatibility
|
||||
image = image.cpu().float()
|
||||
frames = frames.cpu().float()
|
||||
|
||||
# Update batch with decoded image
|
||||
batch.output = image
|
||||
batch.output = frames
|
||||
|
||||
# Offload models if needed
|
||||
if hasattr(self, 'maybe_free_model_hooks'):
|
||||
|
||||
@@ -85,8 +85,9 @@ class DenoisingStage(PipelineStage):
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
|
||||
) # hack
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.SAGE_ATTN_THREE) # hack
|
||||
)
|
||||
|
||||
def forward(
|
||||
@@ -204,8 +205,14 @@ class DenoisingStage(PipelineStage):
|
||||
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
|
||||
|
||||
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
|
||||
if fastvideo_args.boundary_ratio is not None:
|
||||
boundary_timestep = fastvideo_args.boundary_ratio * self.scheduler.num_train_timesteps
|
||||
boundary_ratio = fastvideo_args.pipeline_config.dit_config.boundary_ratio
|
||||
if batch.boundary_ratio is not None:
|
||||
logger.info("Overriding boundary ratio from %s to %s",
|
||||
boundary_ratio, batch.boundary_ratio)
|
||||
boundary_ratio = batch.boundary_ratio
|
||||
|
||||
if boundary_ratio is not None:
|
||||
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
|
||||
else:
|
||||
boundary_timestep = None
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
@@ -247,6 +254,10 @@ class DenoisingStage(PipelineStage):
|
||||
patch_size[2])
|
||||
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
|
||||
|
||||
# Initialize lists for ODE trajectory
|
||||
trajectory_timesteps: list[torch.Tensor] = []
|
||||
trajectory_latents: list[torch.Tensor] = []
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
@@ -427,6 +438,11 @@ class DenoisingStage(PipelineStage):
|
||||
latents = (1. - mask2[0]) * z + mask2[0] * latents
|
||||
# latents = latents.unsqueeze(0)
|
||||
|
||||
# save trajectory latents if needed
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_timesteps.append(t)
|
||||
trajectory_latents.append(latents)
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
(i + 1) > num_warmup_steps and
|
||||
@@ -434,9 +450,28 @@ class DenoisingStage(PipelineStage):
|
||||
and progress_bar is not None):
|
||||
progress_bar.update()
|
||||
|
||||
# Gather results if using sequence parallelism
|
||||
trajectory_tensor: torch.Tensor | None = None
|
||||
if trajectory_latents:
|
||||
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
|
||||
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
|
||||
dim=0)
|
||||
else:
|
||||
trajectory_tensor = None
|
||||
trajectory_timesteps_tensor = None
|
||||
|
||||
# Gather results if using sequence parallelism
|
||||
if sp_group:
|
||||
latents = sequence_model_parallel_all_gather(latents, dim=2)
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_tensor = trajectory_tensor.to(
|
||||
get_local_torch_device())
|
||||
trajectory_tensor = sequence_model_parallel_all_gather(
|
||||
trajectory_tensor, dim=3)
|
||||
|
||||
if trajectory_tensor is not None and trajectory_timesteps_tensor is not None:
|
||||
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
|
||||
batch.trajectory_latents = trajectory_tensor.cpu()
|
||||
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
@@ -115,6 +115,7 @@ class CudaPlatformBase(Platform):
|
||||
|
||||
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND)
|
||||
logger.info("Selected backend: %s", selected_backend)
|
||||
if selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
|
||||
try:
|
||||
from st_attn import sliding_tile_attention # noqa: F401
|
||||
@@ -144,6 +145,20 @@ class CudaPlatformBase(Platform):
|
||||
logger.info(
|
||||
"Sage Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == AttentionBackendEnum.SAGE_ATTN_THREE:
|
||||
try:
|
||||
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell # noqa: F401
|
||||
|
||||
from fastvideo.attention.backends.sage_attn3 import ( # noqa: F401
|
||||
SageAttention3Backend)
|
||||
logger.info("Using Sage Attention 3 backend.")
|
||||
|
||||
return "fastvideo.attention.backends.sage_attn3.SageAttention3Backend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Sage Attention 3 backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
|
||||
try:
|
||||
from vsa import block_sparse_attn # noqa: F401
|
||||
|
||||
@@ -18,6 +18,7 @@ class AttentionBackendEnum(enum.Enum):
|
||||
SLIDING_TILE_ATTN = enum.auto()
|
||||
TORCH_SDPA = enum.auto()
|
||||
SAGE_ATTN = enum.auto()
|
||||
SAGE_ATTN_THREE = enum.auto()
|
||||
VIDEO_SPARSE_ATTN = enum.auto()
|
||||
VMOBA_ATTN = enum.auto()
|
||||
NO_ATTENTION = enum.auto()
|
||||
|
||||
Binary file not shown.
@@ -1 +0,0 @@
|
||||
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":0.50390625,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.19960195198655128,"_runtime":107.325113071}
|
||||
@@ -1 +0,0 @@
|
||||
{"train_loss":0.10545400530099869,"_step":5,"_wandb":{"runtime":35},"step_time":2.2575189135968685,"grad_norm":0.53125,"_runtime":35.701958502,"avg_step_time":2.4754387199878694,"_timestamp":1.7525528420185745e+09,"learning_rate":1e-06,"vsa_sparsity":0,"validation_videos_8_steps":{"videos":[{"caption":"A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.","_type":"video-file","sha256":"ee83e83df073a648f89dcd288cccaed9af765e11fc35c77a6e0cb2ebaa1be5b0","size":475248,"path":"media/videos/validation_videos_8_steps_0_ee83e83df073a648f89d.mp4"},{"size":341490,"path":"media/videos/validation_videos_8_steps_0_d0b2758549d5c82845ca.mp4","caption":"A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","_type":"video-file","sha256":"d0b2758549d5c82845ca3c1ac0db6812877631495e77ac041acdc8ea31f0a5ee"},{"caption":"A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.","_type":"video-file","sha256":"08638381f1607d6ab10684772be38a58b18628fd402a208b25f6af96e765d454","size":436814,"path":"media/videos/validation_videos_8_steps_0_08638381f1607d6ab106.mp4"}],"captions":["A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.","A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table."],"_type":"videos","count":3}}
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import copy
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from abc import abstractmethod
|
||||
@@ -36,10 +37,11 @@ from fastvideo.training.activation_checkpoint import (
|
||||
apply_activation_checkpointing)
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases, count_trainable,
|
||||
EMA_FSDP, clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
get_scheduler, load_distillation_checkpoint, save_distillation_checkpoint,
|
||||
shift_timestep)
|
||||
from fastvideo.utils import is_vsa_available, set_random_seed
|
||||
from fastvideo.utils import (is_vsa_available, maybe_download_model,
|
||||
set_random_seed, verify_model_config_and_directory)
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
@@ -69,18 +71,11 @@ class DistillationPipeline(TrainingPipeline):
|
||||
current_trainstep: int
|
||||
video_latent_shape: tuple[int, ...]
|
||||
video_latent_shape_sp: tuple[int, ...]
|
||||
real_score_transformer: torch.nn.Module
|
||||
fake_score_transformer: torch.nn.Module
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
raise RuntimeError(
|
||||
"create_pipeline_stages should not be called for training pipeline")
|
||||
|
||||
def set_trainable(self) -> None:
|
||||
super().set_trainable()
|
||||
self.modules["real_score_transformer"].requires_grad_(False)
|
||||
self.modules["vae"].requires_grad_(False)
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
"""Initialize the distillation training pipeline with multiple models."""
|
||||
logger.info("Initializing distillation pipeline...")
|
||||
@@ -89,14 +84,35 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
self.noise_scheduler = self.get_module("scheduler")
|
||||
self.vae = self.get_module("vae")
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
self.timestep_shift = self.training_args.pipeline_config.flow_shift
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
|
||||
shift=self.timestep_shift)
|
||||
|
||||
# self.transformer is the generator model
|
||||
self.real_score_transformer = self.get_module("real_score_transformer")
|
||||
self.fake_score_transformer = self.get_module("fake_score_transformer")
|
||||
if training_args.real_score_model_path:
|
||||
logger.info("Loading real score transformer from: %s",
|
||||
training_args.real_score_model_path)
|
||||
self.real_score_transformer = self.load_module_from_path(
|
||||
training_args.real_score_model_path, "transformer",
|
||||
training_args)
|
||||
else:
|
||||
self.real_score_transformer = self.get_module(
|
||||
"real_score_transformer")
|
||||
|
||||
if training_args.fake_score_model_path:
|
||||
logger.info("Loading fake score transformer from: %s",
|
||||
training_args.fake_score_model_path)
|
||||
self.fake_score_transformer = self.load_module_from_path(
|
||||
training_args.fake_score_model_path, "transformer",
|
||||
training_args)
|
||||
else:
|
||||
self.fake_score_transformer = self.get_module(
|
||||
"fake_score_transformer")
|
||||
|
||||
self.real_score_transformer.requires_grad_(False)
|
||||
self.real_score_transformer.eval()
|
||||
self.fake_score_transformer.requires_grad_(True)
|
||||
self.fake_score_transformer.train()
|
||||
|
||||
if training_args.enable_gradient_checkpointing_type is not None:
|
||||
@@ -119,10 +135,13 @@ class DistillationPipeline(TrainingPipeline):
|
||||
if fake_score_lr == 0.0:
|
||||
fake_score_lr = training_args.learning_rate
|
||||
|
||||
betas_str = training_args.fake_score_betas
|
||||
betas = tuple(float(x.strip()) for x in betas_str.split(","))
|
||||
|
||||
self.fake_score_optimizer = torch.optim.AdamW(
|
||||
fake_score_params,
|
||||
lr=fake_score_lr,
|
||||
betas=(0.9, 0.999),
|
||||
betas=betas,
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
@@ -150,8 +169,19 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.training_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
logger.info("Distillation generator model to %s denoising steps",
|
||||
len(self.denoising_step_list))
|
||||
|
||||
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
|
||||
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32))).cuda()
|
||||
self.denoising_step_list = timesteps[1000 -
|
||||
self.denoising_step_list]
|
||||
logger.info("Warping denoising_step_list")
|
||||
|
||||
self.denoising_step_list = self.denoising_step_list.to(
|
||||
get_local_torch_device())
|
||||
logger.info("Distillation generator model to %s denoising steps: %s",
|
||||
len(self.denoising_step_list), self.denoising_step_list)
|
||||
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
|
||||
|
||||
self.min_timestep = int(self.training_args.min_timestep_ratio *
|
||||
@@ -161,6 +191,82 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
|
||||
|
||||
self.generator_ema: EMA_FSDP | None = None
|
||||
if (self.training_args.ema_decay
|
||||
is not None) and (self.training_args.ema_decay > 0.0):
|
||||
self.generator_ema = EMA_FSDP(self.transformer,
|
||||
decay=self.training_args.ema_decay)
|
||||
logger.info("Initialized generator EMA with decay=%s",
|
||||
self.training_args.ema_decay)
|
||||
else:
|
||||
logger.info("Generator EMA disabled (ema_decay <= 0.0)")
|
||||
|
||||
def load_module_from_path(self, model_path: str, module_type: str,
|
||||
training_args: "TrainingArgs"):
|
||||
"""
|
||||
Load a module from a specific path using the same loading logic as the pipeline.
|
||||
|
||||
Args:
|
||||
model_path: Path to the model
|
||||
module_type: Type of module to load (e.g., "transformer")
|
||||
training_args: Training arguments
|
||||
|
||||
Returns:
|
||||
The loaded module
|
||||
"""
|
||||
logger.info("Loading %s from custom path: %s", module_type, model_path)
|
||||
# Set flag to prevent custom weight loading for teacher/critic models
|
||||
training_args._loading_teacher_critic_model = True
|
||||
|
||||
try:
|
||||
from fastvideo.models.loader.component_loader import (
|
||||
PipelineComponentLoader)
|
||||
|
||||
# Download the model if it's a Hugging Face model ID
|
||||
local_model_path = maybe_download_model(model_path)
|
||||
logger.info("Model downloaded/found at: %s", local_model_path)
|
||||
config = verify_model_config_and_directory(local_model_path)
|
||||
|
||||
if module_type not in config:
|
||||
if hasattr(self, '_extra_config_module_map'
|
||||
) and module_type in self._extra_config_module_map:
|
||||
extra_module = self._extra_config_module_map[module_type]
|
||||
if extra_module in config:
|
||||
module_type = extra_module
|
||||
logger.info("Using %s for %s", extra_module,
|
||||
module_type)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Module {module_type} not found in config at {local_model_path}"
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Module {module_type} not found in config at {local_model_path}"
|
||||
)
|
||||
|
||||
module_info = config[module_type]
|
||||
if module_info is None:
|
||||
raise ValueError(
|
||||
f"Module {module_type} has null value in config at {local_model_path}"
|
||||
)
|
||||
|
||||
transformers_or_diffusers, architecture = module_info
|
||||
component_path = os.path.join(local_model_path, module_type)
|
||||
module = PipelineComponentLoader.load_module(
|
||||
module_name=module_type,
|
||||
component_model_path=component_path,
|
||||
transformers_or_diffusers=transformers_or_diffusers,
|
||||
fastvideo_args=training_args,
|
||||
)
|
||||
|
||||
logger.info("Successfully loaded %s from %s", module_type,
|
||||
component_path)
|
||||
return module
|
||||
finally:
|
||||
# Always clean up the flag
|
||||
if hasattr(training_args, '_loading_teacher_critic_model'):
|
||||
delattr(training_args, '_loading_teacher_critic_model')
|
||||
|
||||
@abstractmethod
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
"""Initialize validation pipeline - must be implemented by subclasses."""
|
||||
@@ -170,11 +276,117 @@ class DistillationPipeline(TrainingPipeline):
|
||||
def _prepare_distillation(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Prepare training environment for distillation."""
|
||||
self.transformer.requires_grad_(True)
|
||||
self.transformer.train()
|
||||
self.fake_score_transformer.requires_grad_(True)
|
||||
self.fake_score_transformer.train()
|
||||
|
||||
return training_batch
|
||||
|
||||
def apply_ema_to_model(self, model):
|
||||
"""Apply EMA weights to the model for validation or inference."""
|
||||
if self.generator_ema is not None:
|
||||
with self.generator_ema.apply_to_model(model):
|
||||
return model
|
||||
return model
|
||||
|
||||
def get_ema_model_copy(self) -> torch.nn.Module | None:
|
||||
"""Get a copy of the model with EMA weights applied."""
|
||||
if self.generator_ema is not None:
|
||||
ema_model = copy.deepcopy(self.transformer)
|
||||
self.generator_ema.copy_to_unwrapped(ema_model)
|
||||
return ema_model
|
||||
return None
|
||||
|
||||
def is_ema_ready(self, current_step: int | None = None):
|
||||
"""Check if EMA is ready for use (after ema_start_step)."""
|
||||
if current_step is None:
|
||||
current_step = getattr(self, 'current_trainstep', 0)
|
||||
return (self.generator_ema is not None
|
||||
and current_step >= self.training_args.ema_start_step)
|
||||
|
||||
def save_ema_weights(self, output_dir: str, step: int):
|
||||
"""Save EMA weights separately for inference purposes."""
|
||||
if self.generator_ema is None:
|
||||
logger.warning("Cannot save EMA weights: EMA not initialized")
|
||||
return
|
||||
|
||||
if not self.is_ema_ready():
|
||||
logger.warning(
|
||||
"Cannot save EMA weights: EMA not ready yet (step < ema_start_step)"
|
||||
)
|
||||
return
|
||||
|
||||
try:
|
||||
ema_model = self.get_ema_model_copy()
|
||||
if ema_model is None:
|
||||
logger.warning("Failed to create EMA model copy")
|
||||
return
|
||||
|
||||
ema_save_dir = os.path.join(output_dir, f"ema_checkpoint-{step}")
|
||||
os.makedirs(ema_save_dir, exist_ok=True)
|
||||
|
||||
# save as diffusers format
|
||||
from safetensors.torch import save_file
|
||||
|
||||
from fastvideo.training.training_utils import (
|
||||
custom_to_hf_state_dict, gather_state_dict_on_cpu_rank0)
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(ema_model, device=None)
|
||||
|
||||
if self.global_rank == 0:
|
||||
weight_path = os.path.join(
|
||||
ema_save_dir, "diffusion_pytorch_model.safetensors")
|
||||
diffusers_state_dict = custom_to_hf_state_dict(
|
||||
cpu_state, ema_model.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict, weight_path)
|
||||
|
||||
config_dict = ema_model.hf_config
|
||||
if "dtype" in config_dict:
|
||||
del config_dict["dtype"]
|
||||
config_path = os.path.join(ema_save_dir, "config.json")
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
|
||||
logger.info("EMA weights saved to %s", weight_path)
|
||||
|
||||
del ema_model
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to save EMA weights: %s", str(e))
|
||||
|
||||
def get_ema_stats(self) -> dict[str, Any]:
|
||||
"""Get EMA statistics for monitoring."""
|
||||
if self.generator_ema is None:
|
||||
return {
|
||||
"ema_enabled": False,
|
||||
"ema_decay": None,
|
||||
"ema_start_step": self.training_args.ema_start_step,
|
||||
"ema_ready": False,
|
||||
"ema_step": self.current_trainstep,
|
||||
}
|
||||
|
||||
return {
|
||||
"ema_enabled": True,
|
||||
"ema_decay": self.training_args.ema_decay,
|
||||
"ema_start_step": self.training_args.ema_start_step,
|
||||
"ema_ready": self.is_ema_ready(),
|
||||
"ema_step": self.current_trainstep,
|
||||
}
|
||||
|
||||
def reset_ema(self):
|
||||
"""Reset EMA to current model weights."""
|
||||
if self.generator_ema is not None:
|
||||
logger.info("Resetting EMA to current model weights")
|
||||
self.generator_ema.update(self.transformer)
|
||||
# Force update to current weights by setting decay to 0 temporarily
|
||||
original_decay = self.generator_ema.decay
|
||||
self.generator_ema.decay = 0.0
|
||||
self.generator_ema.update(self.transformer)
|
||||
self.generator_ema.decay = original_decay
|
||||
logger.info("EMA reset completed")
|
||||
else:
|
||||
logger.warning("Cannot reset EMA: EMA not initialized")
|
||||
|
||||
def _build_distill_input_kwargs(
|
||||
self, noise_input: torch.Tensor, timestep: torch.Tensor,
|
||||
text_dict: dict[str, torch.Tensor] | None,
|
||||
@@ -331,6 +543,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
def _dmd_forward(self, generator_pred_video: torch.Tensor,
|
||||
training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""Compute DMD (Diffusion Model Distillation) loss."""
|
||||
original_latent = generator_pred_video
|
||||
with torch.no_grad():
|
||||
timestep = torch.randint(0,
|
||||
self.num_train_timestep, [1],
|
||||
@@ -355,7 +568,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
noisy_latent = self.noise_scheduler.add_noise(
|
||||
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
|
||||
timestep).unflatten(0, (1, generator_pred_video.shape[1]))
|
||||
timestep).detach().unflatten(0,
|
||||
(1, generator_pred_video.shape[1]))
|
||||
|
||||
# fake_score_transformer forward
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
@@ -404,24 +618,24 @@ class DistillationPipeline(TrainingPipeline):
|
||||
pred_real_video_uncond) * self.real_score_guidance_scale
|
||||
|
||||
grad = (faker_score_pred_video - real_score_pred_video) / torch.abs(
|
||||
generator_pred_video - real_score_pred_video).mean()
|
||||
original_latent - real_score_pred_video).mean()
|
||||
grad = torch.nan_to_num(grad)
|
||||
|
||||
dmd_loss = 0.5 * F.mse_loss(
|
||||
generator_pred_video.float(),
|
||||
(generator_pred_video.float() - grad.float()).detach())
|
||||
original_latent.float(),
|
||||
(original_latent.float() - grad.float()).detach())
|
||||
|
||||
training_batch.dmd_latent_vis_dict.update({
|
||||
"training_batch_dmd_fwd_clean_latent":
|
||||
training_batch.latents,
|
||||
"generator_pred_video":
|
||||
generator_pred_video,
|
||||
original_latent.detach(),
|
||||
"real_score_pred_video":
|
||||
real_score_pred_video,
|
||||
real_score_pred_video.detach(),
|
||||
"faker_score_pred_video":
|
||||
faker_score_pred_video,
|
||||
faker_score_pred_video.detach(),
|
||||
"dmd_timestep":
|
||||
timestep,
|
||||
timestep.detach(),
|
||||
})
|
||||
|
||||
return dmd_loss
|
||||
@@ -514,16 +728,17 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"encoder_hidden_states": training_batch.encoder_hidden_states,
|
||||
"encoder_attention_mask": training_batch.encoder_attention_mask,
|
||||
}
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
if getattr(self, "negative_prompt_embeds", None) is not None:
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
|
||||
training_batch.dmd_latent_vis_dict = {}
|
||||
training_batch.fake_score_latent_vis_dict = {}
|
||||
|
||||
training_batch.conditional_dict = conditional_dict
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
training_batch.raw_latent_shape = training_batch.latents.shape
|
||||
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
|
||||
self.video_latent_shape = training_batch.latents.shape
|
||||
@@ -586,8 +801,15 @@ class DistillationPipeline(TrainingPipeline):
|
||||
(dmd_loss / gradient_accumulation_steps).backward()
|
||||
total_dmd_loss += dmd_loss.detach().item()
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer)
|
||||
for param in self.transformer.parameters():
|
||||
# check if the gradient is not None and not zero
|
||||
assert param.grad is not None and param.grad.abs().sum() > 0
|
||||
self.optimizer.step()
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
if self.generator_ema is not None:
|
||||
self.generator_ema.update(self.transformer)
|
||||
|
||||
avg_dmd_loss = torch.tensor(total_dmd_loss /
|
||||
gradient_accumulation_steps,
|
||||
device=self.device)
|
||||
@@ -611,6 +833,9 @@ class DistillationPipeline(TrainingPipeline):
|
||||
fake_score_latent_vis_dict.update(
|
||||
batch_fake.fake_score_latent_vis_dict)
|
||||
self._clip_model_grad_norm_(batch_fake, self.fake_score_transformer)
|
||||
for param in self.fake_score_transformer.parameters():
|
||||
# check if the gradient is not None and not zero
|
||||
assert param.grad is not None and param.grad.abs().sum() > 0
|
||||
self.fake_score_optimizer.step()
|
||||
self.fake_score_lr_scheduler.step()
|
||||
self.lr_scheduler.step()
|
||||
@@ -638,7 +863,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.transformer, self.fake_score_transformer, self.global_rank,
|
||||
self.training_args.resume_from_checkpoint, self.optimizer,
|
||||
self.fake_score_optimizer, self.train_dataloader, self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler, self.noise_random_generator)
|
||||
self.fake_score_lr_scheduler, self.noise_random_generator,
|
||||
self.generator_ema)
|
||||
|
||||
if resumed_step > 0:
|
||||
self.init_steps = resumed_step
|
||||
@@ -669,6 +895,14 @@ class DistillationPipeline(TrainingPipeline):
|
||||
sum(p.numel()
|
||||
for p in self.fake_score_transformer.parameters()) / 1e9)
|
||||
|
||||
if self.generator_ema is not None:
|
||||
logger.info(" Generator EMA enabled with decay: %s",
|
||||
self.training_args.ema_decay)
|
||||
logger.info(" Generator EMA start step: %s",
|
||||
self.training_args.ema_start_step)
|
||||
else:
|
||||
logger.info(" Generator EMA disabled")
|
||||
|
||||
@torch.no_grad()
|
||||
def _log_validation(self, transformer, training_args, global_step) -> None:
|
||||
training_args.inference_mode = True
|
||||
@@ -700,6 +934,18 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
transformer.eval()
|
||||
|
||||
# Optionally use EMA model for validation if available and ready
|
||||
use_ema_for_validation = (self.training_args.use_ema
|
||||
and self.is_ema_ready(global_step))
|
||||
if use_ema_for_validation:
|
||||
logger.info("Using EMA model for validation")
|
||||
validation_transformer = self.transformer
|
||||
ema_context = self.generator_ema.apply_to_model(
|
||||
validation_transformer)
|
||||
else:
|
||||
validation_transformer = transformer
|
||||
ema_context = None
|
||||
|
||||
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]
|
||||
@@ -715,50 +961,98 @@ class DistillationPipeline(TrainingPipeline):
|
||||
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)
|
||||
if ema_context is not None:
|
||||
with ema_context:
|
||||
for validation_batch in validation_dataloader:
|
||||
batch = self._prepare_validation_batch(
|
||||
sampling_param, training_args, validation_batch,
|
||||
num_inference_steps)
|
||||
|
||||
negative_prompt = batch.negative_prompt
|
||||
batch_negative = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt=negative_prompt,
|
||||
prompt_embeds=[],
|
||||
prompt_attention_mask=[],
|
||||
)
|
||||
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
|
||||
batch_negative, training_args)
|
||||
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
negative_prompt = batch.negative_prompt
|
||||
batch_negative = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt=negative_prompt,
|
||||
prompt_embeds=[],
|
||||
prompt_attention_mask=[],
|
||||
)
|
||||
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
|
||||
batch_negative, training_args)
|
||||
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
|
||||
logger.info("rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
|
||||
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)
|
||||
assert batch.prompt is not None and isinstance(
|
||||
batch.prompt, str)
|
||||
step_captions.append(batch.prompt)
|
||||
|
||||
# Run validation inference
|
||||
with torch.no_grad():
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, training_args)
|
||||
samples = output_batch.output
|
||||
if self.rank_in_sp_group != 0:
|
||||
continue
|
||||
# Run validation inference
|
||||
with torch.no_grad():
|
||||
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)
|
||||
# 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)
|
||||
else:
|
||||
# Use original transformer without EMA
|
||||
for validation_batch in validation_dataloader:
|
||||
batch = self._prepare_validation_batch(
|
||||
sampling_param, training_args, validation_batch,
|
||||
num_inference_steps)
|
||||
|
||||
negative_prompt = batch.negative_prompt
|
||||
batch_negative = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt=negative_prompt,
|
||||
prompt_embeds=[],
|
||||
prompt_attention_mask=[],
|
||||
)
|
||||
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
|
||||
batch_negative, training_args)
|
||||
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
|
||||
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
|
||||
with torch.no_grad():
|
||||
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)
|
||||
|
||||
# Log validation results for this step
|
||||
world_group = get_world_group()
|
||||
@@ -835,16 +1129,16 @@ class DistillationPipeline(TrainingPipeline):
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[latent_key] = wandb.Video(
|
||||
video, fps=24, format="mp4") # change to 16 for Wan2.1
|
||||
# Clean up references
|
||||
del video, latents
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[latent_key] = wandb.Video(
|
||||
video, fps=24, format="mp4") # change to 16 for Wan2.1
|
||||
# Clean up references
|
||||
del video, latents
|
||||
|
||||
# Process DMD training data if available - use decode_stage instead of self.vae.decode
|
||||
if 'generator_pred_video' in dmd_latents_vis_dict:
|
||||
@@ -896,14 +1190,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
else:
|
||||
set_random_seed(seed + self.global_rank)
|
||||
|
||||
# Check trainable params
|
||||
num_trainable_generator = round(
|
||||
count_trainable(self.transformer) / 1e9, 3)
|
||||
num_trainable_critic = round(
|
||||
count_trainable(self.fake_score_transformer) / 1e9, 3)
|
||||
logger.info(
|
||||
"rank: %s: # of trainable params in generator: %sB, # of trainable params in critic: %sB",
|
||||
self.global_rank, num_trainable_generator, num_trainable_critic)
|
||||
# Set random seeds for deterministic training
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
@@ -913,6 +1199,10 @@ class DistillationPipeline(TrainingPipeline):
|
||||
device="cpu").manual_seed(self.seed)
|
||||
logger.info("Initialized random seeds with seed: %s", seed)
|
||||
|
||||
# Initialize current_trainstep for EMA ready checks
|
||||
#TODO: check if needed
|
||||
self.current_trainstep = self.init_steps
|
||||
|
||||
# Resume from checkpoint if specified (this will restore random states)
|
||||
if self.training_args.resume_from_checkpoint:
|
||||
self._resume_from_checkpoint()
|
||||
@@ -956,6 +1246,13 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.current_trainstep = step
|
||||
training_batch.current_vsa_sparsity = current_vsa_sparsity
|
||||
|
||||
if (step >= self.training_args.ema_start_step) and \
|
||||
(self.generator_ema is None) and (self.training_args.ema_decay > 0):
|
||||
self.generator_ema = EMA_FSDP(
|
||||
self.transformer, decay=self.training_args.ema_decay)
|
||||
logger.info("Created generator EMA at step %s with decay=%s",
|
||||
step, self.training_args.ema_decay)
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
@@ -969,11 +1266,19 @@ class DistillationPipeline(TrainingPipeline):
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"total_loss": f"{total_loss:.4f}",
|
||||
"generator_loss": f"{generator_loss:.4f}",
|
||||
"fake_score_loss": f"{fake_score_loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
"total_loss":
|
||||
f"{total_loss:.4f}",
|
||||
"generator_loss":
|
||||
f"{generator_loss:.4f}",
|
||||
"fake_score_loss":
|
||||
f"{fake_score_loss:.4f}",
|
||||
"step_time":
|
||||
f"{step_time:.2f}s",
|
||||
"grad_norm":
|
||||
grad_norm,
|
||||
"ema":
|
||||
"✓" if (self.generator_ema is not None and self.is_ema_ready())
|
||||
else "✗",
|
||||
})
|
||||
progress_bar.update(1)
|
||||
|
||||
@@ -1001,6 +1306,15 @@ class DistillationPipeline(TrainingPipeline):
|
||||
if use_vsa:
|
||||
log_data["VSA_train_sparsity"] = current_vsa_sparsity
|
||||
|
||||
if self.generator_ema is not None:
|
||||
log_data["ema_enabled"] = True
|
||||
log_data["ema_decay"] = self.training_args.ema_decay
|
||||
else:
|
||||
log_data["ema_enabled"] = False
|
||||
|
||||
ema_stats = self.get_ema_stats()
|
||||
log_data.update(ema_stats)
|
||||
|
||||
if training_batch.dmd_latent_vis_dict:
|
||||
dmd_additional_logs = {
|
||||
"generator_timestep":
|
||||
@@ -1032,7 +1346,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.global_rank, self.training_args.output_dir, step,
|
||||
self.optimizer, self.fake_score_optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler, self.noise_random_generator)
|
||||
self.fake_score_lr_scheduler, self.noise_random_generator,
|
||||
self.generator_ema)
|
||||
|
||||
if self.transformer:
|
||||
self.transformer.train()
|
||||
@@ -1049,7 +1364,11 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
f"{step}_weight_only",
|
||||
only_save_generator_weight=True)
|
||||
only_save_generator_weight=True,
|
||||
generator_ema=self.generator_ema)
|
||||
|
||||
if self.training_args.use_ema and self.is_ema_ready():
|
||||
self.save_ema_weights(self.training_args.output_dir, step)
|
||||
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
if self.training_args.log_visualization:
|
||||
@@ -1069,7 +1388,11 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.training_args.output_dir, self.training_args.max_train_steps,
|
||||
self.optimizer, self.fake_score_optimizer, self.train_dataloader,
|
||||
self.lr_scheduler, self.fake_score_lr_scheduler,
|
||||
self.noise_random_generator)
|
||||
self.noise_random_generator, self.generator_ema)
|
||||
|
||||
if self.training_args.use_ema and self.is_ema_ready():
|
||||
self.save_ema_weights(self.training_args.output_dir,
|
||||
self.training_args.max_train_steps)
|
||||
|
||||
if get_sp_group():
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
@@ -0,0 +1,406 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from typing import cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
import wandb
|
||||
from fastvideo.dataset.dataloader.schema import (
|
||||
pyarrow_schema_ode_trajectory_text_only)
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
|
||||
SelfForcingFlowMatchScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
|
||||
WanCausalDMDPipeline)
|
||||
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
Training pipeline for ODE-init using precomputed denoising trajectories.
|
||||
|
||||
Supervision: predict the next latent in the stored trajectory by
|
||||
- feeding current latent at timestep t into the transformer to predict noise
|
||||
- stepping the scheduler with the predicted noise
|
||||
- minimizing MSE to the stored next latent at timestep t_next
|
||||
"""
|
||||
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# Match the preprocess/generation scheduler for consistent stepping
|
||||
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
sigma_min=0.0,
|
||||
extra_one_step=True)
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=1000,
|
||||
training=True)
|
||||
|
||||
def set_schemas(self):
|
||||
self.train_dataset_schema = pyarrow_schema_ode_trajectory_text_only
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
super().initialize_training_pipeline(training_args)
|
||||
|
||||
self.noise_scheduler = self.get_module("scheduler")
|
||||
self.vae = self.get_module("vae")
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
self.timestep_shift = self.training_args.pipeline_config.flow_shift
|
||||
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
|
||||
self.noise_scheduler = SelfForcingFlowMatchScheduler(
|
||||
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
|
||||
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
|
||||
training=True)
|
||||
|
||||
logger.info("dmd_denoising_steps: %s",
|
||||
self.training_args.pipeline_config.dmd_denoising_steps)
|
||||
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
|
||||
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32))).cuda()
|
||||
logger.info("timesteps: %s", timesteps)
|
||||
self.dmd_denoising_steps = timesteps[1000 -
|
||||
self.dmd_denoising_steps]
|
||||
logger.info("warped self.dmd_denoising_steps: %s",
|
||||
self.dmd_denoising_steps)
|
||||
else:
|
||||
raise ValueError("warp_denoising_step must be true")
|
||||
|
||||
self.dmd_denoising_steps = self.dmd_denoising_steps.to(
|
||||
get_local_torch_device())
|
||||
|
||||
logger.info("denoising_step_list: %s", self.dmd_denoising_steps)
|
||||
|
||||
logger.info(
|
||||
"Initialized ODE-init training pipeline with %s denoising steps",
|
||||
len(self.dmd_denoising_steps))
|
||||
# Cache for nearest trajectory index per DMD step (computed lazily on first batch)
|
||||
self._cached_closest_idx_per_dmd = None
|
||||
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
|
||||
self.manual_idx = 0
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
args_copy.inference_mode = True
|
||||
# Warm start validation with current transformer
|
||||
self.validation_pipeline = WanCausalDMDPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=args_copy, # type: ignore
|
||||
inference_mode=True,
|
||||
loaded_modules={
|
||||
"transformer": self.get_module("transformer"),
|
||||
},
|
||||
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)
|
||||
|
||||
def _get_next_batch(
|
||||
self,
|
||||
training_batch) -> tuple[TrainingBatch, torch.Tensor, torch.Tensor]:
|
||||
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)
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
batch = next(self.train_loader_iter)
|
||||
|
||||
# Required fields from parquet (ODE trajectory schema)
|
||||
encoder_hidden_states = batch['text_embedding']
|
||||
encoder_attention_mask = batch['text_attention_mask']
|
||||
infos = batch['info_list']
|
||||
|
||||
# Trajectory tensors may include a leading singleton batch dim per row
|
||||
trajectory_latents = batch['trajectory_latents']
|
||||
if trajectory_latents.dim() == 7:
|
||||
# [B, 1, S, C, T, H, W] -> [B, S, C, T, H, W]
|
||||
trajectory_latents = trajectory_latents[:, 0]
|
||||
elif trajectory_latents.dim() == 6:
|
||||
# already [B, S, C, T, H, W]
|
||||
pass
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unexpected trajectory_latents dim: {trajectory_latents.dim()}"
|
||||
)
|
||||
|
||||
trajectory_timesteps = batch['trajectory_timesteps']
|
||||
if trajectory_timesteps.dim() == 3:
|
||||
# [B, 1, S] -> [B, S]
|
||||
trajectory_timesteps = trajectory_timesteps[:, 0]
|
||||
elif trajectory_timesteps.dim() == 2:
|
||||
# [B, S]
|
||||
pass
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unexpected trajectory_timesteps dim: {trajectory_timesteps.dim()}"
|
||||
)
|
||||
# [B, S, C, T, H, W] -> [B, S, T, C, H, W] to match self-forcing
|
||||
trajectory_latents = trajectory_latents.permute(0, 1, 3, 2, 4, 5)
|
||||
|
||||
# Move to device
|
||||
device = get_local_torch_device()
|
||||
training_batch.encoder_hidden_states = encoder_hidden_states.to(
|
||||
device, dtype=torch.bfloat16)
|
||||
training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
device, dtype=torch.bfloat16)
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch, trajectory_latents.to(
|
||||
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
|
||||
|
||||
## TEMP Used for loading the sf .pt files directly
|
||||
"""
|
||||
self.manual_idx = self.manual_idx % 155
|
||||
path = f"/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_pt_vidprom_1000/{self.manual_idx:05d}.pt"
|
||||
logger.info("path: %s", path)
|
||||
self.manual_idx += 1
|
||||
# path = "/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_single_full/00000.pt"
|
||||
b = torch.load(path)
|
||||
training_batch.encoder_hidden_states = b["text_embedding"][0].unsqueeze(
|
||||
0).to(device, dtype=torch.bfloat16)
|
||||
trajectory_latents = b["ode_latent"].to(device, dtype=torch.bfloat16)
|
||||
logger.info("trajectory_latents: %s", trajectory_latents.shape)
|
||||
logger.info("encoder_hidden_states: %s",
|
||||
training_batch.encoder_hidden_states.shape)
|
||||
assert trajectory_latents.shape[1] <= 10, "trajectory_latents.shape[1] must be <= 10"
|
||||
return training_batch, trajectory_latents.to(
|
||||
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
|
||||
"""
|
||||
|
||||
def _get_timestep(self,
|
||||
min_timestep: int,
|
||||
max_timestep: int,
|
||||
batch_size: int,
|
||||
num_frame: int,
|
||||
num_frame_per_block: int,
|
||||
uniform_timestep: bool = False) -> torch.Tensor:
|
||||
if uniform_timestep:
|
||||
timestep = torch.randint(min_timestep,
|
||||
max_timestep, [batch_size, 1],
|
||||
device=self.device,
|
||||
dtype=torch.long).repeat(1, num_frame)
|
||||
return timestep
|
||||
else:
|
||||
timestep = torch.randint(min_timestep,
|
||||
max_timestep, [batch_size, num_frame],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
# logger.info(f"individual timestep: {timestep}")
|
||||
# make the noise level the same within every block
|
||||
timestep = timestep.reshape(timestep.shape[0], -1,
|
||||
num_frame_per_block)
|
||||
timestep[:, :, 1:] = timestep[:, :, 0:1]
|
||||
timestep = timestep.reshape(timestep.shape[0], -1)
|
||||
return timestep
|
||||
|
||||
def _step_predict_next_latent(
|
||||
self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_attention_mask: torch.Tensor
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str,
|
||||
torch.Tensor]]:
|
||||
latent_vis_dict: dict[str, torch.Tensor] = {}
|
||||
device = get_local_torch_device()
|
||||
target_latent = traj_latents[:, -1]
|
||||
|
||||
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
|
||||
B, S, num_frames, num_channels, height, width = traj_latents.shape
|
||||
|
||||
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
|
||||
if self._cached_closest_idx_per_dmd is None:
|
||||
self._cached_closest_idx_per_dmd = torch.tensor(
|
||||
[0, 12, 24, 36], dtype=torch.long).cpu()
|
||||
# [0, 1, 2, 3], dtype=torch.long).cpu()
|
||||
logger.info("self._cached_closest_idx_per_dmd: %s",
|
||||
self._cached_closest_idx_per_dmd)
|
||||
logger.info(
|
||||
"corresponding timesteps: %s", self.noise_scheduler.timesteps[
|
||||
self._cached_closest_idx_per_dmd])
|
||||
|
||||
# Select the K indexes from traj_latents using self._cached_closest_idx_per_dmd
|
||||
# traj_latents: [B, S, C, T, H, W], self._cached_closest_idx_per_dmd: [K]
|
||||
# Output: [B, K, C, T, H, W]
|
||||
assert self._cached_closest_idx_per_dmd is not None
|
||||
relevant_traj_latents = torch.index_select(
|
||||
traj_latents,
|
||||
dim=1,
|
||||
index=self._cached_closest_idx_per_dmd.to(traj_latents.device))
|
||||
logger.info("relevant_traj_latents: %s", relevant_traj_latents.shape)
|
||||
# assert relevant_traj_latents.shape[0] == 1
|
||||
|
||||
indexes = self._get_timestep( # [B, num_frames]
|
||||
0,
|
||||
len(self.dmd_denoising_steps),
|
||||
B,
|
||||
num_frames,
|
||||
3,
|
||||
uniform_timestep=False)
|
||||
logger.info("indexes: %s", indexes.shape)
|
||||
logger.info("indexes: %s", indexes)
|
||||
# noisy_input = relevant_traj_latents[indexes]
|
||||
noisy_input = torch.gather(
|
||||
relevant_traj_latents,
|
||||
dim=1,
|
||||
index=indexes.reshape(B, 1, num_frames, 1, 1,
|
||||
1).expand(-1, -1, -1, num_channels, height,
|
||||
width).to(self.device)).squeeze(1)
|
||||
timestep = self.dmd_denoising_steps[indexes]
|
||||
logger.info("selected timestep for rank %s: %s",
|
||||
self.global_rank,
|
||||
timestep,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Prepare inputs for transformer
|
||||
latent_vis_dict["noisy_input"] = noisy_input.permute(
|
||||
0, 2, 1, 3, 4).detach().clone().cpu()
|
||||
latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3,
|
||||
4).detach().clone().cpu()
|
||||
|
||||
input_kwargs = {
|
||||
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timestep.to(device, dtype=torch.bfloat16),
|
||||
"return_dict": False,
|
||||
}
|
||||
# Predict noise and step the scheduler to obtain next latent
|
||||
with set_forward_context(current_timestep=timestep,
|
||||
attn_metadata=None,
|
||||
forward_batch=None):
|
||||
noise_pred = self.transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=noise_pred.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep.to(dtype=torch.bfloat16).flatten(0, 1),
|
||||
scheduler=self.modules["scheduler"]).unflatten(
|
||||
0, noise_pred.shape[:2])
|
||||
latent_vis_dict["pred_video"] = pred_video.permute(
|
||||
0, 2, 1, 3, 4).detach().clone().cpu()
|
||||
|
||||
return pred_video, target_latent, timestep, latent_vis_dict
|
||||
|
||||
def train_one_step(self, training_batch): # type: ignore[override]
|
||||
self.transformer.train()
|
||||
self.optimizer.zero_grad()
|
||||
training_batch.total_loss = 0.0
|
||||
args = cast(TrainingArgs, self.training_args)
|
||||
|
||||
# Using cached nearest index per DMD step; computation happens in _step_predict_next_latent
|
||||
|
||||
for _ in range(args.gradient_accumulation_steps):
|
||||
training_batch, traj_latents, traj_timesteps = self._get_next_batch(
|
||||
training_batch)
|
||||
text_embeds = training_batch.encoder_hidden_states
|
||||
text_attention_mask = training_batch.encoder_attention_mask
|
||||
assert traj_latents.shape[0] == 1
|
||||
|
||||
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
|
||||
_, S = traj_latents.shape[0], traj_latents.shape[1]
|
||||
if S < 2:
|
||||
raise ValueError("Trajectory must contain at least 2 steps")
|
||||
|
||||
# Forward to predict next latent by stepping scheduler with predicted noise
|
||||
noise_pred, target_latent, t, latent_vis_dict = self._step_predict_next_latent(
|
||||
traj_latents, traj_timesteps, text_embeds, text_attention_mask)
|
||||
|
||||
training_batch.latent_vis_dict.update(latent_vis_dict)
|
||||
|
||||
mask = t != 0
|
||||
|
||||
# Compute loss
|
||||
loss = F.mse_loss(noise_pred[mask],
|
||||
target_latent[mask],
|
||||
reduction="mean")
|
||||
loss = loss / args.gradient_accumulation_steps
|
||||
|
||||
with set_forward_context(current_timestep=t,
|
||||
attn_metadata=None,
|
||||
forward_batch=None):
|
||||
loss.backward()
|
||||
avg_loss = loss.detach().clone()
|
||||
training_batch.total_loss += avg_loss.item()
|
||||
|
||||
# Clip grad and step optimizers
|
||||
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
|
||||
[p for p in self.transformer.parameters() if p.requires_grad],
|
||||
args.max_grad_norm if args.max_grad_norm is not None else 0.0)
|
||||
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
if grad_norm is None:
|
||||
grad_value = 0.0
|
||||
else:
|
||||
try:
|
||||
if isinstance(grad_norm, torch.Tensor):
|
||||
grad_value = float(grad_norm.detach().float().item())
|
||||
else:
|
||||
grad_value = float(grad_norm)
|
||||
except Exception:
|
||||
grad_value = 0.0
|
||||
training_batch.grad_norm = grad_value
|
||||
return training_batch
|
||||
|
||||
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
|
||||
training_args: TrainingArgs, step: int):
|
||||
"""Add visualization data to wandb logging and save frames to disk."""
|
||||
wandb_loss_dict = {}
|
||||
latents_vis_dict = training_batch.latent_vis_dict
|
||||
latent_log_keys = ['noisy_input', 'x0', 'pred_video']
|
||||
for latent_key in latent_log_keys:
|
||||
assert latent_key in latents_vis_dict and latents_vis_dict[
|
||||
latent_key] is not None
|
||||
latent = latents_vis_dict[latent_key]
|
||||
pixel_latent = self.validation_pipeline.decoding_stage.decode(
|
||||
latent, training_args)
|
||||
|
||||
video = pixel_latent.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[latent_key] = wandb.Video(
|
||||
video, fps=16, format="mp4") # change to 16 for Wan2.1
|
||||
# Clean up references
|
||||
del video, pixel_latent, latent
|
||||
|
||||
# Log to wandb
|
||||
if self.global_rank == 0:
|
||||
wandb.log(wandb_loss_dict, step=step)
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting ODE-init training pipeline...")
|
||||
pipeline = ODEInitTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("ODE-init training pipeline done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.fastvideo_args import TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.dit_cpu_offload = False
|
||||
main(args)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -478,7 +478,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
current_vsa_sparsity = current_decay_times * vsa_decay_rate
|
||||
elif vmoba_available:
|
||||
# TODO: add vmoba sparsity scheduling here
|
||||
pass
|
||||
current_vsa_sparsity = 0.0
|
||||
else:
|
||||
current_vsa_sparsity = 0.0
|
||||
|
||||
@@ -520,6 +520,10 @@ 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:
|
||||
if self.training_args.log_visualization:
|
||||
self.visualize_intermediate_latents(training_batch,
|
||||
self.training_args,
|
||||
step)
|
||||
self._log_validation(self.transformer, self.training_args, step)
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
trainable_params = round(
|
||||
@@ -720,3 +724,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# Re-enable gradients for training
|
||||
training_args.inference_mode = False
|
||||
transformer.train()
|
||||
|
||||
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
|
||||
training_args: TrainingArgs, step: int):
|
||||
"""Add visualization data to wandb logging and save frames to disk."""
|
||||
raise NotImplementedError(
|
||||
"Visualize intermediate latents is not implemented for training pipeline"
|
||||
)
|
||||
|
||||
@@ -202,6 +202,7 @@ def save_distillation_checkpoint(generator_transformer,
|
||||
generator_scheduler=None,
|
||||
fake_score_scheduler=None,
|
||||
noise_generator=None,
|
||||
generator_ema=None,
|
||||
only_save_generator_weight=False) -> None:
|
||||
"""
|
||||
Save distillation checkpoint with both generator and fake_score models.
|
||||
@@ -233,6 +234,8 @@ def save_distillation_checkpoint(generator_transformer,
|
||||
if generator_scheduler is not None:
|
||||
generator_states["scheduler"] = SchedulerWrapper(
|
||||
generator_scheduler)
|
||||
if generator_ema is not None:
|
||||
generator_states["ema"] = generator_ema.state_dict()
|
||||
|
||||
generator_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
|
||||
"generator")
|
||||
@@ -402,7 +405,8 @@ def load_distillation_checkpoint(generator_transformer,
|
||||
dataloader=None,
|
||||
generator_scheduler=None,
|
||||
fake_score_scheduler=None,
|
||||
noise_generator=None) -> int:
|
||||
noise_generator=None,
|
||||
generator_ema=None) -> int:
|
||||
"""
|
||||
Load distillation checkpoint with both generator and fake_score models.
|
||||
Returns the step number from which training should resume.
|
||||
@@ -456,6 +460,20 @@ def load_distillation_checkpoint(generator_transformer,
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Load EMA state if available and generator_ema is provided
|
||||
if generator_ema is not None:
|
||||
try:
|
||||
ema_state = generator_states.get("ema")
|
||||
if ema_state is not None:
|
||||
generator_ema.load_state_dict(ema_state)
|
||||
logger.info("rank: %s, generator EMA state loaded successfully",
|
||||
rank)
|
||||
else:
|
||||
logger.info("rank: %s, no EMA state found in checkpoint", rank)
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed to load EMA state: %s", rank,
|
||||
str(e))
|
||||
|
||||
# Load critic distributed checkpoint
|
||||
critic_dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint",
|
||||
"critic")
|
||||
@@ -1282,3 +1300,169 @@ def get_scheduler(
|
||||
|
||||
def count_trainable(model: torch.nn.Module) -> int:
|
||||
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
|
||||
|
||||
class EMA_FSDP:
|
||||
"""
|
||||
FSDP2-friendly EMA with two modes:
|
||||
- mode="local_shard" (default): maintain float32 CPU EMA of local parameter shards on every rank.
|
||||
Provides a context manager to temporarily swap EMA weights into the live model for teacher forward.
|
||||
- mode="rank0_full": maintain a consolidated float32 CPU EMA of full parameters on rank 0 only
|
||||
using gather_state_dict_on_cpu_rank0(). Useful for checkpoint export; not for teacher forward.
|
||||
|
||||
Usage (local_shard for CM teacher):
|
||||
ema = EMA_FSDP(model, decay=0.999, mode="local_shard")
|
||||
for step in ...:
|
||||
ema.update(model)
|
||||
with ema.apply_to_model(model):
|
||||
with torch.no_grad():
|
||||
y_teacher = model(...)
|
||||
|
||||
Usage (rank0_full for export):
|
||||
ema = EMA_FSDP(model, decay=0.999, mode="rank0_full")
|
||||
ema.update(model)
|
||||
ema.state_dict() # on rank 0
|
||||
"""
|
||||
|
||||
def __init__(self, module, decay: float = 0.999, mode: str = "local_shard"):
|
||||
self.decay = float(decay)
|
||||
self.mode = mode
|
||||
self.shadow: dict[str, torch.Tensor] = {}
|
||||
self.rank = dist.get_rank() if dist.is_initialized() else 0
|
||||
if self.mode not in {"local_shard", "rank0_full"}:
|
||||
raise ValueError(f"Unsupported EMA_FSDP mode: {self.mode}")
|
||||
self._init_shadow(module)
|
||||
|
||||
@staticmethod
|
||||
def _to_local_tensor(t: torch.Tensor) -> torch.Tensor:
|
||||
# DTensor-aware to_local fetch; fall back to raw tensor
|
||||
try:
|
||||
from torch.distributed.tensor import DTensor # type: ignore
|
||||
if isinstance(t, DTensor):
|
||||
return t.to_local()
|
||||
except Exception:
|
||||
pass
|
||||
return t
|
||||
|
||||
@torch.no_grad()
|
||||
def _init_shadow(self, module):
|
||||
if self.mode == "rank0_full":
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(module, device=None)
|
||||
if self.rank == 0:
|
||||
self.shadow = {
|
||||
k: v.detach().clone().float().cpu()
|
||||
for k, v in cpu_state.items()
|
||||
}
|
||||
else:
|
||||
self.shadow = {}
|
||||
return
|
||||
|
||||
# local_shard: maintain EMA of local shards for requires_grad params
|
||||
self.shadow = {}
|
||||
for name, p in module.named_parameters():
|
||||
if not p.requires_grad:
|
||||
continue
|
||||
local = self._to_local_tensor(p.detach())
|
||||
self.shadow[name] = local.clone().float().cpu()
|
||||
|
||||
@torch.no_grad()
|
||||
def update(self, module):
|
||||
d = self.decay
|
||||
if self.mode == "rank0_full":
|
||||
if self.rank != 0:
|
||||
return
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(module, device=None)
|
||||
for n, v in cpu_state.items():
|
||||
v_cpu = v.detach().float().cpu()
|
||||
if n not in self.shadow:
|
||||
self.shadow[n] = v_cpu.clone()
|
||||
else:
|
||||
self.shadow[n].mul_(d).add_(v_cpu, alpha=1.0 - d)
|
||||
return
|
||||
|
||||
# local_shard: update local shard EMA on every rank
|
||||
for name, p in module.named_parameters():
|
||||
if not p.requires_grad:
|
||||
continue
|
||||
local = self._to_local_tensor(p.detach())
|
||||
v_cpu = local.float().cpu()
|
||||
if name not in self.shadow:
|
||||
self.shadow[name] = v_cpu.clone()
|
||||
else:
|
||||
self.shadow[name].mul_(d).add_(v_cpu, alpha=1.0 - d)
|
||||
|
||||
def state_dict(self) -> dict[str, torch.Tensor]:
|
||||
if self.mode == "rank0_full":
|
||||
return {
|
||||
k: v.clone()
|
||||
for k, v in self.shadow.items()
|
||||
} if self.rank == 0 else {}
|
||||
return {k: v.clone() for k, v in self.shadow.items()}
|
||||
|
||||
def load_state_dict(self, sd: dict[str, torch.Tensor]):
|
||||
self.shadow = {k: v.clone() for k, v in sd.items()}
|
||||
|
||||
@torch.no_grad()
|
||||
def copy_to_unwrapped(self, module) -> None:
|
||||
"""
|
||||
Copy EMA weights into a non-sharded (unwrapped) module. Intended for export/eval.
|
||||
For mode="rank0_full", only rank 0 has the full EMA state.
|
||||
"""
|
||||
if self.mode == "rank0_full" and self.rank != 0:
|
||||
return
|
||||
name_to_param = dict(module.named_parameters())
|
||||
for n, w in self.shadow.items():
|
||||
if n in name_to_param:
|
||||
p = name_to_param[n]
|
||||
p.data.copy_(w.to(dtype=p.dtype, device=p.device))
|
||||
|
||||
class _ApplyEMACtx:
|
||||
|
||||
def __init__(self, ema: "EMA_FSDP", module):
|
||||
self.ema = ema
|
||||
self.module = module
|
||||
self.saved: dict[str, torch.Tensor] = {}
|
||||
|
||||
def __enter__(self):
|
||||
if self.ema.mode != "local_shard":
|
||||
raise RuntimeError(
|
||||
"EMA apply_to_model is only supported for mode='local_shard'"
|
||||
)
|
||||
with torch.no_grad():
|
||||
for name, p in self.module.named_parameters():
|
||||
if not p.requires_grad:
|
||||
continue
|
||||
# Save local shard
|
||||
p_local = EMA_FSDP._to_local_tensor(p.detach())
|
||||
if p_local.numel() == 0:
|
||||
# Nothing to swap on this rank for this param
|
||||
continue
|
||||
self.saved[name] = p_local.clone().to(device=p_local.device,
|
||||
dtype=p_local.dtype)
|
||||
if name in self.ema.shadow:
|
||||
ema_cpu = self.ema.shadow[name]
|
||||
if ema_cpu.numel() != p_local.numel():
|
||||
# Shard shape mismatch (e.g., empty shard here), skip
|
||||
continue
|
||||
# Copy EMA shard into local param shard
|
||||
p_local.copy_(
|
||||
ema_cpu.to(dtype=p_local.dtype,
|
||||
device=p_local.device))
|
||||
return self.module
|
||||
|
||||
def __exit__(self, exc_type, exc, tb):
|
||||
with torch.no_grad():
|
||||
for name, p in self.module.named_parameters():
|
||||
if name in self.saved:
|
||||
p_local = EMA_FSDP._to_local_tensor(p.detach())
|
||||
if p_local.numel() == 0:
|
||||
continue
|
||||
saved_local = self.saved[name]
|
||||
if saved_local.numel() != p_local.numel():
|
||||
continue
|
||||
p_local.copy_(saved_local)
|
||||
self.saved.clear()
|
||||
return False
|
||||
|
||||
def apply_to_model(self, module):
|
||||
return EMA_FSDP._ApplyEMACtx(self, module)
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
|
||||
WanCausalDMDPipeline)
|
||||
from fastvideo.training.self_forcing_distillation_pipeline import (
|
||||
SelfForcingDistillationPipeline)
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanSelfForcingDistillationPipeline(SelfForcingDistillationPipeline):
|
||||
"""
|
||||
A self-forcing distillation pipeline for Wan that uses the self-forcing methodology
|
||||
with DMD for video generation.
|
||||
"""
|
||||
_required_config_modules = [
|
||||
"scheduler", "transformer", "vae", "real_score_transformer",
|
||||
"fake_score_transformer"
|
||||
]
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
args_copy.inference_mode = True
|
||||
validation_pipeline = WanCausalDMDPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=args_copy, # type: ignore
|
||||
inference_mode=True,
|
||||
loaded_modules={"transformer": self.get_module("transformer")},
|
||||
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)
|
||||
|
||||
self.validation_pipeline = validation_pipeline
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting Wan self-forcing distillation pipeline...")
|
||||
|
||||
pipeline = WanSelfForcingDistillationPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("Wan self-forcing distillation pipeline completed")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.fastvideo_args import TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -23,10 +23,14 @@ from typing import Any, TypeVar, cast
|
||||
|
||||
import cloudpickle
|
||||
import filelock
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
import yaml
|
||||
from diffusers.loaders.lora_base import (
|
||||
_best_guess_weight_name) # watch out for potetential removal from diffusers
|
||||
from einops import rearrange
|
||||
from huggingface_hub import snapshot_download
|
||||
from remote_pdb import RemotePdb
|
||||
from torch.distributed.fsdp import MixedPrecisionPolicy
|
||||
@@ -886,3 +890,17 @@ def best_output_size(w, h, dw, dh, expected_area):
|
||||
return ow1, oh1
|
||||
else:
|
||||
return ow2, oh2
|
||||
|
||||
|
||||
def save_decoded_latents_as_video(decoded_latents: list[torch.Tensor],
|
||||
output_path: str, fps: int):
|
||||
# Process outputs
|
||||
videos = rearrange(decoded_latents, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
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))
|
||||
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
imageio.mimsave(output_path, frames, fps=fps, format="mp4")
|
||||
|
||||
@@ -1,20 +1,19 @@
|
||||
import dataclasses
|
||||
import gc
|
||||
import multiprocessing
|
||||
import os
|
||||
import random
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
from datasets import Dataset, Video, load_dataset
|
||||
|
||||
from fastvideo.configs.configs import (DatasetType, PreprocessConfig,
|
||||
VideoLoaderType)
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.distributed.parallel_state import get_world_rank, get_world_size
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
|
||||
@@ -155,31 +154,17 @@ class VideoForwardBatchBuilder:
|
||||
|
||||
|
||||
class ParquetDatasetSaver:
|
||||
"""Component for saving and writing Parquet datasets"""
|
||||
"""Component for saving and writing Parquet datasets using shared parquet_io."""
|
||||
|
||||
def __init__(self,
|
||||
flush_frequency: int,
|
||||
samples_per_file: int,
|
||||
schema_fields: list[str],
|
||||
record_creator: Callable[..., list[dict[str, Any]]],
|
||||
file_writer_fn: Callable | None = None):
|
||||
"""
|
||||
Initialize ParquetDatasetSaver
|
||||
|
||||
Args:
|
||||
schema_fields: schema fields list
|
||||
record_creator: Function for creating records
|
||||
file_writer_fn: Function for writing records to files, uses default implementation if None
|
||||
"""
|
||||
def __init__(self, flush_frequency: int, samples_per_file: int,
|
||||
schema: pa.Schema,
|
||||
record_creator: Callable[..., list[dict[str, Any]]]):
|
||||
self.flush_frequency = flush_frequency
|
||||
self.samples_per_file = samples_per_file
|
||||
self.schema_fields = schema_fields
|
||||
self.schema = schema
|
||||
self.create_records_from_batch = record_creator
|
||||
self.file_writer_fn: Callable[
|
||||
[tuple], int] = file_writer_fn or self._default_file_writer_fn
|
||||
self.all_tables: list[pa.Table] = []
|
||||
self.num_processed_samples: int = 0
|
||||
self.num_saved_files: int = 0
|
||||
self._writer: ParquetDatasetWriter | None = None
|
||||
|
||||
def save_and_write_parquet_batch(
|
||||
self,
|
||||
@@ -227,16 +212,17 @@ class ParquetDatasetSaver:
|
||||
|
||||
if batch_data:
|
||||
self.num_processed_samples += len(batch_data)
|
||||
# Convert batch data to PyArrow arrays
|
||||
table = self._convert_batch_to_pyarrow_table(batch_data)
|
||||
|
||||
# Store the table in a list for later processing
|
||||
self.all_tables.append(table)
|
||||
table = records_to_table(batch_data, self.schema)
|
||||
if self._writer is None:
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
self._writer = ParquetDatasetWriter(
|
||||
out_dir=output_dir, samples_per_file=self.samples_per_file)
|
||||
self._writer.append_table(table)
|
||||
logger.debug("Collected batch with %s samples", len(table))
|
||||
|
||||
# If flush is needed
|
||||
if self.num_processed_samples >= self.flush_frequency:
|
||||
self.flush_tables(output_dir)
|
||||
self.flush_tables()
|
||||
|
||||
def _process_non_padded_embeddings(
|
||||
self, prompt_embeds: torch.Tensor,
|
||||
@@ -259,146 +245,31 @@ class ParquetDatasetSaver:
|
||||
|
||||
return non_padded_embeds
|
||||
|
||||
def _convert_batch_to_pyarrow_table(self,
|
||||
batch_data: list[dict]) -> pa.Table:
|
||||
"""Convert batch data to PyArrow table"""
|
||||
arrays = []
|
||||
def flush_tables(self, write_remainder: bool = False):
|
||||
"""Flush buffered records to disk.
|
||||
|
||||
for field in self.schema_fields:
|
||||
if field.endswith('_bytes'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.binary()))
|
||||
elif field.endswith('_shape'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.list_(pa.int32())))
|
||||
elif field in ['width', 'height', 'num_frames']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.int32()))
|
||||
elif field in ['duration_sec', 'fps']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.float32()))
|
||||
else:
|
||||
arrays.append(pa.array([record[field]
|
||||
for record in batch_data]))
|
||||
|
||||
return pa.Table.from_arrays(arrays, names=self.schema_fields)
|
||||
|
||||
def flush_tables(self, output_dir: str):
|
||||
"""Flush collected tables to disk"""
|
||||
if not hasattr(self, 'all_tables') or not self.all_tables:
|
||||
Args:
|
||||
output_dir: Directory where parquet files are written. Kept for API
|
||||
symmetry (writer already configured with this path).
|
||||
write_remainder: If True, also write any leftover rows smaller than
|
||||
``samples_per_file`` as a final small file. Useful for the last flush.
|
||||
"""
|
||||
if self._writer is None:
|
||||
return
|
||||
|
||||
logger.debug("Combining %d batches...", len(self.all_tables))
|
||||
combined_table = pa.concat_tables(self.all_tables)
|
||||
assert len(combined_table) == self.num_processed_samples
|
||||
logger.debug("Total samples collected: %d", len(combined_table))
|
||||
|
||||
# Calculate total number of chunks needed, putting remainder into self.all_tables
|
||||
total_files = max(self.num_processed_samples // self.samples_per_file,
|
||||
1)
|
||||
|
||||
logger.debug("Fixed samples per parquet file: %d",
|
||||
self.samples_per_file)
|
||||
logger.debug("Total number of parquet files: %d", total_files)
|
||||
logger.debug(
|
||||
"Total samples to be processed: %d (putting %d samples into self.all_tables)",
|
||||
total_files * self.samples_per_file,
|
||||
self.num_processed_samples % self.samples_per_file)
|
||||
|
||||
# Split work among processes
|
||||
num_workers = int(min(multiprocessing.cpu_count(), total_files))
|
||||
files_per_worker = (total_files + num_workers - 1) // num_workers
|
||||
|
||||
logger.debug("Using %d workers to process %d files", num_workers,
|
||||
total_files)
|
||||
logger.debug("Files per worker: %s", files_per_worker)
|
||||
|
||||
# Prepare work ranges
|
||||
work_ranges = []
|
||||
for i in range(num_workers):
|
||||
start_idx = i * files_per_worker
|
||||
end_idx = min((i + 1) * files_per_worker, total_files)
|
||||
if start_idx < total_files:
|
||||
work_ranges.append((start_idx, end_idx, combined_table, i,
|
||||
output_dir, self.samples_per_file))
|
||||
|
||||
total_written = 0
|
||||
failed_ranges = []
|
||||
with ProcessPoolExecutor(max_workers=num_workers) as executor:
|
||||
futures = {
|
||||
executor.submit(self.file_writer_fn, work_range): work_range
|
||||
for work_range in work_ranges
|
||||
}
|
||||
for future in futures:
|
||||
try:
|
||||
written = future.result()
|
||||
total_written += written
|
||||
logger.info("Processed file with %s samples", written)
|
||||
except Exception as e:
|
||||
work_range = futures[future]
|
||||
failed_ranges.append(work_range)
|
||||
logger.error("Failed to process range %s-%s: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
# Retry failed ranges sequentially
|
||||
if failed_ranges:
|
||||
logger.warning("Retrying %s failed ranges sequentially",
|
||||
len(failed_ranges))
|
||||
for work_range in failed_ranges:
|
||||
try:
|
||||
total_written += self.file_writer_fn(work_range)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to process range %s-%s after retry: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
self.num_saved_files += total_files
|
||||
|
||||
# Clear tables list
|
||||
self.all_tables = []
|
||||
if self.num_processed_samples > self.samples_per_file:
|
||||
saved_samples = total_files * self.samples_per_file
|
||||
self.all_tables.append(combined_table.slice(saved_samples))
|
||||
self.num_processed_samples -= saved_samples
|
||||
else:
|
||||
self.num_processed_samples = 0
|
||||
|
||||
del combined_table
|
||||
gc.collect()
|
||||
_ = self._writer.flush(write_remainder=write_remainder)
|
||||
# Reset processed sample count modulo samples_per_file
|
||||
remainder = self.num_processed_samples % self.samples_per_file
|
||||
self.num_processed_samples = 0 if write_remainder else remainder
|
||||
|
||||
def clean_up(self) -> None:
|
||||
"""Clean up all tables"""
|
||||
self.all_tables = []
|
||||
self.flush_tables(write_remainder=True)
|
||||
self._writer = None
|
||||
self.num_processed_samples = 0
|
||||
self.num_saved_files = 0
|
||||
gc.collect()
|
||||
|
||||
def _default_file_writer_fn(self, args_tuple: tuple) -> int:
|
||||
"""Default chunk processing implementation"""
|
||||
start_idx, end_idx, combined_table, worker_id, output_dir, samples_per_file = args_tuple
|
||||
|
||||
written_count = 0
|
||||
for file_idx in range(start_idx, end_idx):
|
||||
start_row = file_idx * samples_per_file
|
||||
end_row = min(start_row + samples_per_file, len(combined_table))
|
||||
|
||||
if start_row >= len(combined_table):
|
||||
break
|
||||
|
||||
chunk_table = combined_table.slice(start_row, end_row - start_row)
|
||||
|
||||
# Write to file
|
||||
output_file = os.path.join(
|
||||
output_dir,
|
||||
f"chunk_{file_idx + self.num_saved_files:06d}.parquet")
|
||||
pq.write_table(chunk_table, output_file)
|
||||
written_count += len(chunk_table)
|
||||
|
||||
return written_count
|
||||
def __del__(self):
|
||||
self.clean_up()
|
||||
|
||||
|
||||
def build_dataset(preprocess_config: PreprocessConfig, split: str,
|
||||
|
||||
@@ -4,6 +4,8 @@ from typing import cast
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from fastvideo.configs.configs import PreprocessConfig
|
||||
from fastvideo.dataset.dataloader.record_schema import (
|
||||
basic_t2v_record_creator, i2v_record_creator)
|
||||
from fastvideo.dataset.dataloader.schema import (pyarrow_schema_i2v,
|
||||
pyarrow_schema_t2v)
|
||||
from fastvideo.distributed.parallel_state import get_world_rank
|
||||
@@ -13,8 +15,6 @@ from fastvideo.pipelines.pipeline_registry import PipelineType
|
||||
from fastvideo.workflow.preprocess.components import (
|
||||
ParquetDatasetSaver, PreprocessingDataValidator, VideoForwardBatchBuilder,
|
||||
build_dataset)
|
||||
from fastvideo.workflow.preprocess.record_schema import (
|
||||
basic_t2v_record_creator, i2v_record_creator)
|
||||
from fastvideo.workflow.workflow_base import WorkflowBase
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -85,16 +85,16 @@ class PreprocessWorkflow(WorkflowBase):
|
||||
# record creator
|
||||
if self.fastvideo_args.workload_type == WorkloadType.I2V:
|
||||
record_creator = i2v_record_creator
|
||||
schema_fields = [f.name for f in pyarrow_schema_i2v]
|
||||
schema = pyarrow_schema_i2v
|
||||
else:
|
||||
record_creator = basic_t2v_record_creator
|
||||
schema_fields = [f.name for f in pyarrow_schema_t2v]
|
||||
schema = pyarrow_schema_t2v
|
||||
processed_dataset_saver = ParquetDatasetSaver(
|
||||
flush_frequency=self.fastvideo_args.preprocess_config.
|
||||
flush_frequency,
|
||||
samples_per_file=self.fastvideo_args.preprocess_config.
|
||||
samples_per_file,
|
||||
schema_fields=schema_fields,
|
||||
schema=schema,
|
||||
record_creator=record_creator,
|
||||
)
|
||||
self.add_component("processed_dataset_saver", processed_dataset_saver)
|
||||
|
||||
@@ -35,8 +35,7 @@ class PreprocessWorkflowI2V(PreprocessWorkflow):
|
||||
self.processed_dataset_saver.save_and_write_parquet_batch(
|
||||
forward_batch, self.training_dataset_output_dir)
|
||||
|
||||
self.processed_dataset_saver.flush_tables(
|
||||
self.training_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables()
|
||||
self.processed_dataset_saver.clean_up()
|
||||
|
||||
# Validation dataset preprocessing
|
||||
@@ -51,6 +50,5 @@ class PreprocessWorkflowI2V(PreprocessWorkflow):
|
||||
|
||||
self.processed_dataset_saver.save_and_write_parquet_batch(
|
||||
forward_batch, self.validation_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables(
|
||||
self.validation_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables()
|
||||
self.processed_dataset_saver.clean_up()
|
||||
|
||||
@@ -35,8 +35,7 @@ class PreprocessWorkflowT2V(PreprocessWorkflow):
|
||||
self.processed_dataset_saver.save_and_write_parquet_batch(
|
||||
forward_batch, self.training_dataset_output_dir)
|
||||
|
||||
self.processed_dataset_saver.flush_tables(
|
||||
self.training_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables()
|
||||
self.processed_dataset_saver.clean_up()
|
||||
|
||||
# Validation dataset preprocessing
|
||||
@@ -51,6 +50,5 @@ class PreprocessWorkflowT2V(PreprocessWorkflow):
|
||||
|
||||
self.processed_dataset_saver.save_and_write_parquet_batch(
|
||||
forward_batch, self.validation_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables(
|
||||
self.validation_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables()
|
||||
self.processed_dataset_saver.clean_up()
|
||||
|
||||
+1
-1
@@ -34,7 +34,7 @@ dependencies = [
|
||||
|
||||
# Miscellaneous Utilities
|
||||
"tqdm", "pytest", "PyYAML==6.0.1", "protobuf>=5.28.3",
|
||||
"gradio>=5.22.0", "moviepy>=2.0.0", "flask",
|
||||
"gradio==5.41.0", "moviepy>=2.0.0", "flask",
|
||||
"flask_restful", "aiohttp", "huggingface_hub", "cloudpickle",
|
||||
# System & Monitoring Tools
|
||||
"gpustat", "watch", "remote-pdb",
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
#!/bin/bash
|
||||
|
||||
# Create output directory if it doesn't exist
|
||||
mkdir -p preprocess_output_text
|
||||
|
||||
# Launch 8 jobs, one for each node
|
||||
# Each node processes 8 consecutive files (64 total files / 8 nodes = 8 files per node)
|
||||
for node_id in {0..7}; do
|
||||
# Calculate the starting file number for this node
|
||||
start_file=$((node_id * 8 + 1))
|
||||
|
||||
echo "Launching text-only node $node_id with files v2m_${start_file}.txt to v2m_$((start_file + 7)).txt"
|
||||
echo "sbatch --job-name=text-${node_id} --output=preprocess_output_text/preprocess-text-node-${node_id}.out --error=preprocess_output_text/preprocess-text-node-${node_id}.err scripts/preprocess/syn_text.slurm $start_file $node_id"
|
||||
|
||||
sbatch --job-name=text-${node_id} \
|
||||
--output=preprocess_output_text/preprocess-text-node-${node_id}.out \
|
||||
--error=preprocess_output_text/preprocess-text-node-${node_id}.err \
|
||||
scripts/preprocess/syn_text.slurm $start_file $node_id
|
||||
done
|
||||
|
||||
echo "All 8 text-only nodes launched successfully!"
|
||||
@@ -0,0 +1,88 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks-per-node=8
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=16
|
||||
#SBATCH --mem=960G
|
||||
#SBATCH --exclusive
|
||||
#SBATCH --time=72:00:00
|
||||
|
||||
# conda init
|
||||
# source ~/conda/miniconda/bin/activate
|
||||
# PYTHON_VIRTUAL_ENVIRONMENT=fastvideo-train-yq
|
||||
# conda activate $PYTHON_VIRTUAL_ENVIRONMENT
|
||||
nvidia-smi
|
||||
nvidia-smi --query-compute-apps=pid,process_name,used_memory --format=csv
|
||||
|
||||
echo " "
|
||||
echo " Number of nodes:= " $SLURM_JOB_NUM_NODES
|
||||
echo " GPUs per node:= " $SLURM_JOB_GPUS
|
||||
echo " Running on multiple nodes/GPU devices for TEXT-ONLY preprocessing"
|
||||
echo ""
|
||||
echo " Run started at:- "
|
||||
date
|
||||
|
||||
# Accept parameters from launch script
|
||||
START_FILE=${1:-1} # Starting file number for this node
|
||||
NODE_ID=${2:-0} # Node identifier (0-7)
|
||||
|
||||
num_gpus=1
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
# Start port number - we'll increment for each job
|
||||
base_port=$((29603 + NODE_ID * 100)) # Different port range per node
|
||||
|
||||
# Create an array of CUDA device IDs
|
||||
gpu_ids=(0 1 2 3 4 5 6 7)
|
||||
|
||||
GPU_NUM=1
|
||||
MODEL_TYPE="wan"
|
||||
|
||||
echo "NODE_ID: $NODE_ID"
|
||||
echo "START_FILE: $START_FILE"
|
||||
echo "Base port for this node: $base_port"
|
||||
echo "Processing TEXT-ONLY data"
|
||||
|
||||
# Run 8 parallel preprocessing jobs on this node
|
||||
for i in {1..8}; do
|
||||
# Calculate port for this job
|
||||
port=$((base_port + i))
|
||||
|
||||
# Get GPU ID using modulo to cycle through available GPUs
|
||||
gpu=${gpu_ids[((i-1))]}
|
||||
|
||||
# Calculate which file this GPU should process
|
||||
file_num=$((START_FILE + i - 1))
|
||||
DATA_MERGE_PATH="prompts/v2m_${file_num}.txt"
|
||||
|
||||
# Create unique output directory based on node and GPU for text-only processing
|
||||
OUTPUT_DIR="data/test-text-preprocessing/Node_${NODE_ID}_GPU_${i}_File_${file_num}"
|
||||
|
||||
start_cpu=$(( (i-1)*2 )) # Reduced CPU allocation for 8 nodes
|
||||
end_cpu=$(( start_cpu+1 ))
|
||||
|
||||
echo "Starting GPU $gpu processing text-only file v2m_${file_num}.txt on port $port, output: $OUTPUT_DIR"
|
||||
|
||||
# Run the text-only preprocessing command in background
|
||||
CUDA_VISIBLE_DEVICES=$gpu taskset -c ${start_cpu}-${end_cpu} torchrun --nnodes=1 --nproc_per_node=$GPU_NUM --master_port $port \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_BASE \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 2 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "text_only" &
|
||||
done
|
||||
|
||||
# Wait for all jobs on this node to complete
|
||||
wait
|
||||
|
||||
echo "All text-only processing blocks completed!"
|
||||
@@ -21,4 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--validation_dataset_file $VALIDATION_PATH \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--preprocess_task "i2v"
|
||||
--preprocess_task "i2v"
|
||||
|
||||
@@ -22,4 +22,4 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--samples_per_file 1 \
|
||||
--flush_frequency 1 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v"
|
||||
--preprocess_task "t2v"
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
from fastvideo.dataset.dataloader.parquet_io import (
|
||||
ParquetDatasetWriter,
|
||||
records_to_table,
|
||||
)
|
||||
|
||||
|
||||
def test_records_to_table_types():
|
||||
schema = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
pa.field("vae_latent_bytes", pa.binary()),
|
||||
pa.field("vae_latent_shape", pa.list_(pa.int64())),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("width", pa.int64()),
|
||||
])
|
||||
records = [{
|
||||
"id": "a",
|
||||
"vae_latent_bytes": b"\x00\x01",
|
||||
"vae_latent_shape": [1, 2, 3],
|
||||
"duration_sec": 1.5,
|
||||
"width": 640,
|
||||
}]
|
||||
|
||||
table = records_to_table(records, schema)
|
||||
assert table.schema == schema
|
||||
assert table.num_rows == 1
|
||||
cols = {name: table.column(name).to_pylist()[0] for name in schema.names}
|
||||
assert cols["id"] == "a"
|
||||
assert isinstance(cols["vae_latent_bytes"], (bytes, bytearray))
|
||||
assert cols["vae_latent_shape"] == [1, 2, 3]
|
||||
assert abs(cols["duration_sec"] - 1.5) < 1e-6
|
||||
assert cols["width"] == 640
|
||||
|
||||
|
||||
def test_writer_flush_and_remainder(tmp_path: Path):
|
||||
schema = pa.schema([pa.field("id", pa.string())])
|
||||
records = [{"id": str(i)} for i in range(25)]
|
||||
table = records_to_table(records, schema)
|
||||
|
||||
out_dir = tmp_path / "out"
|
||||
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
|
||||
writer.append_table(table)
|
||||
written = writer.flush(num_workers=1)
|
||||
assert written == 20
|
||||
|
||||
files = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files) == 2
|
||||
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
|
||||
assert total_rows == 20
|
||||
|
||||
# Append remainder to complete another chunk
|
||||
extra = records_to_table([{"id": str(i)} for i in range(5)], schema)
|
||||
writer.append_table(extra)
|
||||
written2 = writer.flush(num_workers=1)
|
||||
assert written2 == 10
|
||||
files2 = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files2) == 3
|
||||
total_rows2 = sum(pq.read_table(str(f)).num_rows for f in files2)
|
||||
assert total_rows2 == 30
|
||||
|
||||
|
||||
def test_writer_flush_write_remainder(tmp_path: Path):
|
||||
schema = pa.schema([pa.field("id", pa.string())])
|
||||
# 25 rows, 10 per file => 2 full files + 1 remainder(5)
|
||||
records = [{"id": str(i)} for i in range(25)]
|
||||
table = records_to_table(records, schema)
|
||||
|
||||
out_dir = tmp_path / "out_last"
|
||||
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
|
||||
writer.append_table(table)
|
||||
# First flush writes 20
|
||||
written1 = writer.flush(num_workers=1)
|
||||
assert written1 == 20
|
||||
# Final flush with remainder
|
||||
written2 = writer.flush(num_workers=1, write_remainder=True)
|
||||
assert written2 == 5
|
||||
files = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files) == 3
|
||||
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
|
||||
assert total_rows == 25
|
||||
|
||||
|
||||
def test_writer_parallel_workers(tmp_path: Path):
|
||||
schema = pa.schema([pa.field("id", pa.string())])
|
||||
# 40 rows, 10 per file => 4 files
|
||||
records = [{"id": str(i)} for i in range(40)]
|
||||
table = records_to_table(records, schema)
|
||||
|
||||
out_dir = tmp_path / "out_parallel"
|
||||
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
|
||||
writer.append_table(table)
|
||||
written = writer.flush(num_workers=2)
|
||||
assert written == 40
|
||||
|
||||
# Ensure files exist under worker subdirs
|
||||
worker_dirs = [p for p in out_dir.iterdir() if p.is_dir() and p.name.startswith("worker_")]
|
||||
assert len(worker_dirs) >= 1
|
||||
files = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files) == 4
|
||||
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
|
||||
assert total_rows == 40
|
||||
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.dataset.dataloader.record_schema import (
|
||||
basic_t2v_record_creator,
|
||||
i2v_record_creator,
|
||||
ode_text_only_record_creator,
|
||||
text_only_record_creator,
|
||||
)
|
||||
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
|
||||
|
||||
|
||||
def _mk_basic_batch(N: int) -> PreprocessBatch:
|
||||
batch = PreprocessBatch(data_type="video")
|
||||
batch.video_file_name = [f"vid_{i}" for i in range(N)]
|
||||
batch.prompt = [f"caption_{i}" for i in range(N)]
|
||||
batch.width = [640 for _ in range(N)]
|
||||
batch.height = [360 for _ in range(N)]
|
||||
batch.fps = [4 for _ in range(N)]
|
||||
batch.num_frames = [2 for _ in range(N)]
|
||||
# Latents: shape (N, C, T, H, W); per-record use latents[idx]
|
||||
batch.latents = np.zeros((N, 4, 2, 8, 8), dtype=np.float32)
|
||||
# Prompt embeds: list of per-record arrays [Seq, Dim]
|
||||
batch.prompt_embeds = [np.ones((6, 16), dtype=np.float32) for _ in range(N)]
|
||||
return batch
|
||||
|
||||
|
||||
def test_basic_t2v_record_creator_fields():
|
||||
N = 2
|
||||
batch = _mk_basic_batch(N)
|
||||
|
||||
records = basic_t2v_record_creator(batch)
|
||||
assert isinstance(records, list) and len(records) == N
|
||||
|
||||
for i, rec in enumerate(records):
|
||||
assert rec["id"] == batch.video_file_name[i]
|
||||
# Latents bytes/shape/dtype
|
||||
assert isinstance(rec["vae_latent_bytes"], (bytes, bytearray))
|
||||
assert rec["vae_latent_shape"] == list(batch.latents[i].shape)
|
||||
assert rec["vae_latent_dtype"] == str(batch.latents[i].dtype)
|
||||
# Text embedding
|
||||
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
|
||||
assert rec["text_embedding_shape"] == list(batch.prompt_embeds[i].shape)
|
||||
assert rec["text_embedding_dtype"] == str(batch.prompt_embeds[i].dtype)
|
||||
# Meta
|
||||
assert rec["caption"] == batch.prompt[i]
|
||||
assert rec["media_type"] == "video"
|
||||
assert rec["width"] == int(batch.width[i])
|
||||
assert rec["height"] == int(batch.height[i])
|
||||
assert rec["num_frames"] == batch.latents[i].shape[1]
|
||||
|
||||
|
||||
def test_i2v_record_creator_additional_fields():
|
||||
N = 3
|
||||
batch = _mk_basic_batch(N)
|
||||
# image_embeds is a list of length 1, with an array of shape [N, D]
|
||||
batch.image_embeds = [np.ones((N, 32), dtype=np.float32)]
|
||||
# first frame latent per record
|
||||
batch.image_latent = np.zeros((N, 4, 1, 8, 8), dtype=np.float32)
|
||||
# pil image per record
|
||||
batch.pil_image = np.zeros((N, 8, 8, 3), dtype=np.uint8)
|
||||
|
||||
records = i2v_record_creator(batch)
|
||||
assert isinstance(records, list) and len(records) == N
|
||||
|
||||
for i, rec in enumerate(records):
|
||||
# clip feature
|
||||
assert isinstance(rec["clip_feature_bytes"], (bytes, bytearray))
|
||||
assert rec["clip_feature_shape"] == list(batch.image_embeds[0][i].shape)
|
||||
assert rec["clip_feature_dtype"] == str(batch.image_embeds[0][i].dtype)
|
||||
# first frame latent
|
||||
assert isinstance(rec["first_frame_latent_bytes"], (bytes, bytearray))
|
||||
assert rec["first_frame_latent_shape"] == list(batch.image_latent[i].shape)
|
||||
assert rec["first_frame_latent_dtype"] == str(batch.image_latent[i].dtype)
|
||||
# pil image
|
||||
assert isinstance(rec["pil_image_bytes"], (bytes, bytearray))
|
||||
assert rec["pil_image_shape"] == list(batch.pil_image[i].shape)
|
||||
assert rec["pil_image_dtype"] == str(batch.pil_image[i].dtype)
|
||||
|
||||
|
||||
def test_ode_text_only_record_creator():
|
||||
video_name = "ex"
|
||||
caption = "a prompt"
|
||||
text_embedding = np.ones((6, 16), dtype=np.float32)
|
||||
traj = np.ones((5, 4, 2, 2), dtype=np.float32)
|
||||
tsteps = np.arange(5, dtype=np.float32)
|
||||
|
||||
rec = ode_text_only_record_creator(
|
||||
video_name=video_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=caption,
|
||||
trajectory_latents=traj,
|
||||
trajectory_timesteps=tsteps,
|
||||
)
|
||||
assert rec["id"] == f"text_{video_name}"
|
||||
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
|
||||
assert rec["text_embedding_shape"] == list(text_embedding.shape)
|
||||
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
|
||||
assert rec["file_name"] == video_name
|
||||
assert rec["caption"] == caption
|
||||
assert rec["media_type"] == "text"
|
||||
# Trajectory fields
|
||||
assert isinstance(rec["trajectory_latents_bytes"], (bytes, bytearray))
|
||||
assert rec["trajectory_latents_shape"] == list(traj.shape)
|
||||
assert rec["trajectory_latents_dtype"] == str(traj.dtype)
|
||||
assert isinstance(rec["trajectory_timesteps_bytes"], (bytes, bytearray))
|
||||
assert rec["trajectory_timesteps_shape"] == list(tsteps.shape)
|
||||
assert rec["trajectory_timesteps_dtype"] == str(tsteps.dtype)
|
||||
|
||||
|
||||
def test_text_only_record_creator():
|
||||
text_name = "note1"
|
||||
caption = "a prompt"
|
||||
text_embedding = np.ones((7, 16), dtype=np.float32)
|
||||
rec = text_only_record_creator(
|
||||
text_name=text_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=caption,
|
||||
)
|
||||
assert rec["id"] == f"text_{text_name}"
|
||||
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
|
||||
assert rec["text_embedding_shape"] == list(text_embedding.shape)
|
||||
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
|
||||
assert rec["caption"] == caption
|
||||
+1
-2
@@ -6,7 +6,7 @@ import pytest
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
|
||||
from tests.utils import compute_video_ssim_torchvision, write_ssim_results
|
||||
from diffusers import DiffusionPipeline
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines import build_pipeline
|
||||
@@ -192,4 +192,3 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
|
||||
|
||||
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user