Compare commits

...
Author SHA1 Message Date
SolitaryThinker 48f9690c47 load_video 2025-06-24 10:46:20 -07:00
SolitaryThinker 32df3bcca3 update i2v script 2025-06-24 03:20:09 -07:00
SolitaryThinker b8c81191b6 improve script format 2025-06-24 03:01:55 -07:00
SolitaryThinker e506c4074f remove print and enable first val 2025-06-24 01:55:05 -07:00
SolitaryThinker 1b3914da84 update scripts 2025-06-22 21:09:30 -07:00
SolitaryThinker f471fd3f02 i2v working 2025-06-22 20:59:37 -07:00
SolitaryThinker 2e6d5c5304 t2v working again 2025-06-22 19:45:44 -07:00
SolitaryThinker 5db34184c7 t2v example 2025-06-22 21:53:05 +00:00
SolitaryThinker 37252bf62c f 2025-06-22 11:58:18 +00:00
SolitaryThinker e9263f7d2b update 2025-06-22 04:20:34 -07:00
SolitaryThinker 93afb86c20 update 2025-06-22 04:19:10 -07:00
SolitaryThinker 0d944ba9c1 slrm 2025-06-22 04:09:09 -07:00
SolitaryThinker 3d78604281 update path 2025-06-22 03:54:39 -07:00
SolitaryThinker 0694b0c5eb exmaple scripts 2025-06-22 03:44:57 -07:00
SolitaryThinker b6c5644d40 cleanup 2025-06-22 03:07:37 -07:00
SolitaryThinker a9089fa358 fix pil image 2025-06-22 02:29:51 -07:00
SolitaryThinker 6079a98fd7 i2v preprocess 2025-06-21 19:18:39 -07:00
SolitaryThinker 65c0fcb633 checkpoint 2025-06-21 18:02:43 -07:00
SolitaryThinker 82e3641264 checkpoint 2025-06-21 18:00:29 -07:00
William Lin 8741d204a5 [Training] Refactor and improve validation datasets (#539) 2025-06-21 17:58:35 -07:00
Wenxuan Tan cdc85f58a8 [chore] Bump torch to 2.7.1 to support Blackwell (#483) 2025-06-20 22:10:56 -07:00
William Lin 0262d2f089 [misc] [training] Reorganize training pipeline (#533) 2025-06-20 20:42:25 -07:00
William Lin 62c0343465 [bugfix] [VSA] Fix layernorm type for VSA Wan2.1 TransformerBlock (#534) 2025-06-20 00:24:51 -07:00
William Lin 1e1a023fb0 [bugfix] Fix stage validator for multi text encoder models (#535) 2025-06-19 22:49:16 -07:00
William Lin 1d2517ad8e [misc] Remove gradient checking code (#532) 2025-06-18 23:29:25 -07:00
William Lin d41186cb4a [Feat] Add Stage input and output verification (#523) 2025-06-18 23:29:11 -07:00
78e0c7eec9 Specify cu128 Pytorch installation (#530)
Co-authored-by: Edenzzzz <wtan45@wisc.edu>
Co-authored-by: Wenxuan Tan <wenxuan.tan@wisc.edu>
2025-06-18 20:02:50 -05:00
Wenxuan Tan 1c41a94b62 [Refactor] Move dict_to_3d_list under utils (#507) 2025-06-18 13:34:37 -07:00
Yongqi Chen 2e66aafe20 [Bugfix][Readme]Fix readme website bugs and add VSA finetune docs (#531) 2025-06-17 22:48:29 -07:00
Yongqi ChenandWill Lin 55074bda76 [CI] Add STA-inference/VSA-training test (#527)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-06-17 21:13:06 -07:00
William Lin de65bec2b7 [Ci] add sta and vsa install to docker image (#528) 2025-06-17 18:09:48 -07:00
Yongqi Chen 7664dd0de3 [Bugfix][Inference]Fix envs.attn_backend (#525) 2025-06-17 18:38:06 -05:00
William Linandkevin314 019a88ced4 [CI][bugfix] Use new 3.12 docker image (#526)
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-06-17 15:37:08 -07:00
Kevin Lin 72de11abcc [CI] Add current PR test workflow to Buildkite/Modal (#512) 2025-06-17 13:29:22 -07:00
Kevin Lin d71a4ebffc [CI] Update Docker image to flash-attn 2.8.0 / CUDA 12.8 (#524) 2025-06-16 17:48:23 -07:00
William Lin 1089ab43bf [bugfix] [Training] use diffusers fp32layernorm for wan2.1 (#490) 2025-06-15 22:45:48 -07:00
105 changed files with 4733 additions and 1379 deletions
+66
View File
@@ -0,0 +1,66 @@
env:
IMAGE_VERSION: "py3.12-latest"
steps:
- block: "Start Build"
blocked_state: "running"
prompt: "Approve build?"
- label: "Trigger Tests"
command: |
echo "Current working directory: $(pwd)"
echo "Current branch:"
git branch --show-current
echo "Full diff:"
git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD
plugins:
- monorepo-diff#v1.4.0:
diff: "git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD"
watch:
- path:
- "fastvideo/v1/models/encoders/**"
- "fastvideo/v1/models/loaders/**"
- "fastvideo/v1/tests/encoders/**"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=encoder
agents:
queue: "default"
- path:
- "fastvideo/v1/models/vaes/**"
- "fastvideo/v1/models/loaders/**"
- "fastvideo/v1/tests/vaes/**"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=vae
agents:
queue: "default"
- path:
- "fastvideo/v1/models/dits/**"
- "fastvideo/v1/models/loaders/**"
- "fastvideo/v1/tests/transformers/**"
- "fastvideo/v1/layers/**"
- "fastvideo/v1/attention/**"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Transformer Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=transformer
agents:
queue: "default"
- path: "fastvideo/v1/**/*.py"
config:
command: "timeout 60m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=ssim
agents:
queue: "default"
+91
View File
@@ -0,0 +1,91 @@
#!/bin/bash
set -uo pipefail
log() {
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
}
log "=== Starting Modal test execution ==="
# Change to the project directory
cd "$(dirname "$0")/../.."
PROJECT_ROOT=$(pwd)
log "Project root: $PROJECT_ROOT"
# Install Modal if not available
if ! python3 -m modal --version &> /dev/null; then
log "Modal not found, installing..."
python3 -m pip install modal
# Verify installation
if ! python3 -m modal --version &> /dev/null; then
log "Error: Failed to install modal. Please install it manually."
exit 1
fi
fi
log "modal version: $(python3 -m modal --version)"
# Set up Modal authentication using Buildkite secrets
log "Setting up Modal authentication from Buildkite secrets..."
MODAL_TOKEN_ID=$(buildkite-agent secret get modal_token_id)
MODAL_TOKEN_SECRET=$(buildkite-agent secret get modal_token_secret)
if [ -n "$MODAL_TOKEN_ID" ] && [ -n "$MODAL_TOKEN_SECRET" ]; then
log "Retrieved Modal credentials from Buildkite secrets"
python3 -m modal token set --token-id "$MODAL_TOKEN_ID" --token-secret "$MODAL_TOKEN_SECRET" --profile buildkite-ci --activate --verify
if [ $? -eq 0 ]; then
log "Modal authentication successful"
else
log "Error: Failed to set Modal credentials"
exit 1
fi
else
log "Error: Could not retrieve Modal credentials from Buildkite secrets."
log "Please ensure 'modal_token_id' and 'modal_token_secret' secrets are set in Buildkite."
exit 1
fi
MODAL_TEST_FILE="fastvideo/v1/tests/modal/pr_test.py"
if [ -z "${TEST_TYPE:-}" ]; then
log "Error: TEST_TYPE environment variable is not set"
exit 1
fi
log "Test type: $TEST_TYPE"
case "$TEST_TYPE" in
"encoder")
log "Running encoder tests..."
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
;;
"vae")
log "Running VAE tests..."
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
;;
"transformer")
log "Running transformer tests..."
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
;;
"ssim")
log "Running SSIM tests..."
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
exit 1
;;
esac
log "Executing: $MODAL_COMMAND"
eval "$MODAL_COMMAND"
TEST_EXIT_CODE=$?
if [ $TEST_EXIT_CODE -eq 0 ]; then
log "Modal test completed successfully"
else
log "Error: Modal test failed with exit code: $TEST_EXIT_CODE"
fi
log "=== Test execution completed with exit code: $TEST_EXIT_CODE ==="
exit $TEST_EXIT_CODE
+1 -2
View File
@@ -160,8 +160,7 @@ def execute_command(pod_id):
setup_steps = [
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
f"cd /workspace/{repo_name}",
"source /opt/conda/etc/profile.d/conda.sh",
"conda activate fastvideo-dev",
"source $HOME/.local/bin/env && source /opt/venv/bin/activate",
args.test_command
]
+83 -13
View File
@@ -12,12 +12,14 @@ on:
paths:
- "fastvideo/**/*.py"
- ".github/workflows/pr-test.yml"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
workflow_dispatch:
inputs:
custom_image:
description: "Custom image from this repository (default: fastvideo-dev:latest)"
description: "Custom image from this repository (default: fastvideo-dev:py3.12-latest)"
required: false
default: "fastvideo-dev:latest"
default: "fastvideo-dev:py3.12-latest"
type: string
run_encoder_test:
description: "Run encoder-test"
@@ -44,6 +46,16 @@ on:
required: false
default: false
type: boolean
run_training_test_VSA:
description: "Run training-test-VSA"
required: false
default: false
type: boolean
run_inference_test_STA:
description: "Run inference-test-STA"
required: false
default: false
type: boolean
run_nightly_test:
description: "Run nightly-test"
required: false
@@ -70,6 +82,8 @@ jobs:
vae-test: ${{ steps.filter.outputs.vae-test }}
transformer-test: ${{ steps.filter.outputs.transformer-test }}
training-test: ${{ steps.filter.outputs.training-test }}
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
steps:
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
@@ -80,18 +94,34 @@ jobs:
- 'fastvideo/v1/models/encoders/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/encoders/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
vae-test:
- 'fastvideo/v1/models/vaes/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/vaes/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
transformer-test:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/transformers/**'
- 'fastvideo/v1/layers/**'
- 'fastvideo/v1/attention/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
training-test:
- 'fastvideo/v1/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
training-test-VSA:
- 'fastvideo/v1/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
inference-test-STA:
- 'fastvideo/v1/**'
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
encoder-test:
needs: change-filter
@@ -104,8 +134,8 @@ jobs:
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -122,8 +152,8 @@ jobs:
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -140,8 +170,8 @@ jobs:
gpu_type: "NVIDIA L40S"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -168,7 +198,7 @@ jobs:
volume_size: 200
disk_size: 200
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
timeout_minutes: 60
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -186,8 +216,48 @@ jobs:
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/training -srP"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/Vanilla -srP"
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 }}
training-test-VSA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "training-test-VSA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/VSA -srP"
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 }}
inference-test-STA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "inference-test-STA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/inference/STA -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -203,8 +273,8 @@ jobs:
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
+2 -2
View File
@@ -91,7 +91,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
## Distillation and Finetuning
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetuning.html)
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html)
## 📑 Development Plan
@@ -111,7 +111,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
## 🤝 Contributing
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/developer_guide/overview.html)
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
## Acknowledgement
We learned and reused code from the following projects:
+44 -20
View File
@@ -1,7 +1,9 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
@@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
rm Miniconda3-latest-Linux-x86_64.sh
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
ENV PATH=/opt/conda/bin:$PATH
# Set CUDA environment variables
ENV CUDA_HOME=/usr/local/cuda-12.8
ENV PATH=${CUDA_HOME}/bin:${PATH}
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
RUN conda create --name fastvideo-dev python=3.12.9 -y
SHELL ["/bin/bash", "-c"]
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
@@ -27,22 +37,36 @@ COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
conda clean -afy
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.12 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.0.post2 --no-build-isolation
COPY . .
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[dev] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Remove authentication headers
RUN git config --unset-all http.https://github.com/.extraheader || true
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
git submodule update --init --recursive && \
python setup_sta.py install
# Set up automatic conda environment activation for all shells
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
# Ensure .bashrc is sourced for SSH login shells
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
git submodule update --init --recursive && \
python setup_vsa.py install
EXPOSE 22
+7 -1
View File
@@ -1,7 +1,7 @@
(sta-demo)=
# 🔍 Demo
There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
<div style="text-align: center;">
<video controls width="800">
@@ -9,3 +9,9 @@ There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
Your browser does not support the video tag.
</video>
</div>
You can run STA using the following command:
```bash
bash scripts/inference/v1_inference_wan_STA.sh
```
+7
View File
@@ -16,6 +16,13 @@ bash scripts/finetune/finetune_mochi.sh # for mochi
```
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
## ⚡ Finetune with VSA
Follow [data_preprocess.md](#v0-data-preprocess) to get parquet files for preproccessed latent, and then run:
```bash
bash scripts/finetune/finetune_v1_VSA.sh
```
## ⚡ Lora Finetune
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
@@ -5,7 +5,7 @@ export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
base_port=29503
num_gpu=$(nvidia-smi --query-gpu=gpu_name --format=csv,noheader | wc -l)
num_gpu=1
gpu_ids=$(seq 0 $((num_gpu-1)))
skip_time_steps=12
@@ -14,7 +14,7 @@ STA_mode="STA_searching"
for i in $gpu_ids; do
port=$((base_port+i))
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
--prompt_path ./assets/prompt_extend_${i}.txt \
--prompt_path ./assets/prompt_${i}.txt \
--output_path $output_path \
--STA_mode $STA_mode &
sleep 1
@@ -27,7 +27,7 @@ STA_mode="STA_tuning"
for i in $gpu_ids; do
port=$((base_port+i))
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
--prompt_path ./assets/prompt_extend_${i}.txt \
--prompt_path ./assets/prompt_${i}.txt \
--output_path $output_path \
--STA_mode $STA_mode \
--skip_time_steps $skip_time_steps &
@@ -0,0 +1,88 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_i2v/validation_parquet_dataset/"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_i2v_finetune"
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 8
--tp_size 8
--hsdp_replicate_dim 1
--hsdp_shard_dim 8
)
# Model arguments
model_args=(
--model_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
--pretrained_model_name_or_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_preprocessed_path "$VALIDATION_DIR"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 6000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--cfg 0.0
--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/v1/training/wan_i2v_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,97 @@
#!/bin/bash
#SBATCH --job-name=FV_2N_14B
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[400-550]
#SBATCH --mem=1440G
#SBATCH --output=4n_i2v/4n_i2v_%j.out
#SBATCH --error=4n_i2v/4n_i2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv
# 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_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR=data/crush-smol_processed_i2v/combined_parquet_dataset
VALIDATION_DIR=data/crush-smol_processed_i2v/validation_parquet_dataset
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# 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/v1/training/wan_i2v_training_pipeline.py\
--model_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
--cache_dir "/home/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_preprocessed_path "$VALIDATION_DIR"\
--train_batch_size=1\
--num_latent_t 16 \
--num_gpus $NUM_GPUS \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES \
--hsdp_shard_dim $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 10\
--gradient_accumulation_steps=2\
--max_train_steps=10000 \
--learning_rate=5e-5\
--mixed_precision="bf16"\
--checkpointing_steps=11000 \
--validation_steps 100\
--validation_sampling_steps "40" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"\
--tracker_project_name wan_i2v_finetune \
--num_height 480 \
--num_width 832 \
--num_frames 77 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 1e-4 \
--not_apply_cfg_solver \
--dit_precision "fp32" \
--max_grad_norm 1.0
@@ -0,0 +1,24 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_i2v/"
VALIDATION_PATH="examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation.json"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "i2v"
@@ -0,0 +1,31 @@
{
"data": [
{
"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.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"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.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -0,0 +1,88 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=4,5
# Training arguments
training_args=(
--tracker_project_name "wan_t2v_finetune"
--output_dir "outputs/wan_t2v_finetune"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--hsdp_replicate_dim 1 \
--hsdp_shard_dim $NUM_GPUS \
)
# 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_preprocessed_path $VALIDATION_DIR
--validation_steps 100
--validation_sampling_steps "50"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 6000
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--cfg 0.0
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,98 @@
#!/bin/bash
#SBATCH --job-name=FV_2N_14B
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[400-550]
#SBATCH --mem=1440G
#SBATCH --output=4n_i2v/4n_i2v_%j.out
#SBATCH --error=4n_i2v/4n_i2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv
# 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_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# 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/v1/training/wan_training_pipeline.py\
--model_path $MODEL_PATH \
--inference_mode False\
--pretrained_model_name_or_path $MODEL_PATH \
--cache_dir "/home/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_preprocessed_path "$VALIDATION_DIR"\
--train_batch_size=1\
--num_latent_t 8 \
--num_gpus $NUM_GPUS \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES \
--hsdp_shard_dim $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 10\
--gradient_accumulation_steps=1\
--max_train_steps=10000 \
--learning_rate=5e-5\
--mixed_precision="bf16"\
--checkpointing_steps=11000 \
--validation_steps 100\
--validation_sampling_steps "40" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"\
--tracker_project_name wan_i2v_finetune \
--num_height 480 \
--num_width 832 \
--num_frames 77 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 1e-4 \
--not_apply_cfg_solver \
--dit_precision "fp32" \
--max_grad_norm 1.0
@@ -0,0 +1,13 @@
{
"data": [
{
"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.",
"image_path": null,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -0,0 +1,24 @@
# export WANDB_MODE="offline"
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/"
VALIDATION_PATH="examples/training/finetune/wan_t2v_1_3b/crush_smol/validation.json"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "t2v"
@@ -0,0 +1,31 @@
{
"data": [
{
"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.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"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.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
+2 -15
View File
@@ -6,6 +6,8 @@ from typing import Any, Dict, List, Optional, Tuple
import numpy as np
from fastvideo.v1.utils import dict_to_3d_list
def configure_sta(mode: str = 'STA_searching',
layer_num: int = 40,
@@ -349,21 +351,6 @@ def select_best_mask_strategy(
return best_mask_strategy, overall_sparsity, strategy_counts
def dict_to_3d_list(mask_strategy: Optional[Dict[str, List[int]]],
t_max: int = 50,
l_max: int = 60,
h_max: int = 24) -> List[List[List[Optional[List[int]]]]]:
result: List[List[List[Optional[List[int]]]]] = [[[
None for _ in range(h_max)
] for _ in range(l_max)] for _ in range(t_max)]
if mask_strategy is None:
return result
for key, value in mask_strategy.items():
t, layer_idx, h = map(int, key.split('_'))
result[t][layer_idx][h] = value
return result
def save_mask_search_results(
mask_search_final_result: List[Dict[str, List[float]]],
prompt: str,
+1 -1
View File
@@ -10,8 +10,8 @@ from fastvideo.v1.attention.selector import get_attn_backend
__all__ = [
"DistributedAttention",
"DistributedAttention_VSA",
"LocalAttention",
"DistributedAttention_VSA",
"AttentionBackend",
"AttentionMetadata",
"AttentionMetadataBuilder",
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import json
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Type
from typing import Any, List, Optional, Type
import torch
from einops import rearrange
@@ -17,33 +17,11 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.utils import dict_to_3d_list
logger = init_logger(__name__)
# TODO(will-refactor): move this to a utils file
def dict_to_3d_list(
mask_strategy: Dict[str,
Any]) -> List[List[List[Optional[torch.Tensor]]]]:
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
max_timesteps_idx = max(
timesteps_idx for timesteps_idx, layer_idx, head_idx in indices) + 1
max_layer_idx = max(layer_idx
for timesteps_idx, layer_idx, head_idx in indices) + 1
max_head_idx = max(head_idx
for timesteps_idx, layer_idx, head_idx in indices) + 1
result = [[[None for _ in range(max_head_idx)]
for _ in range(max_layer_idx)] for _ in range(max_timesteps_idx)]
for key, value in mask_strategy.items():
timesteps_idx, layer_idx, head_idx = map(int, key.split('_'))
result[timesteps_idx][layer_idx][head_idx] = value
return result
class RangeDict(dict):
def __getitem__(self, item: int) -> str:
+11 -1
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import json
from dataclasses import asdict, dataclass, field, fields
from enum import Enum
from typing import Any, Callable, Dict, List, Optional, Tuple, Union, cast
import torch
@@ -16,6 +17,15 @@ from fastvideo.v1.utils import (FlexibleArgumentParser, StoreBoolean,
logger = init_logger(__name__)
class STA_Mode(str, Enum):
"""STA (Sliding Tile Attention) modes."""
STA_INFERENCE = "STA_inference"
STA_SEARCHING = "STA_searching"
STA_TUNING = "STA_tuning"
STA_TUNING_CFG = "STA_tuning_cfg"
NONE = None
def preprocess_text(prompt: str) -> str:
return prompt
@@ -76,7 +86,7 @@ class PipelineConfig:
# STA (Sliding Tile Attention) parameters
mask_strategy_file_path: Optional[str] = None
STA_mode: Optional[str] = None
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
skip_time_steps: int = 15
# Compilation
+17 -20
View File
@@ -1,19 +1,17 @@
import os
# SPDX-License-Identifier: Apache-2.0
from torchvision import transforms
from torchvision.transforms import Lambda
from transformers import AutoTokenizer
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
from fastvideo.v1.dataset.parquet_dataset_map_style import (
build_parquet_map_style_dataloader)
from fastvideo.v1.dataset.preprocessing_datasets import (
VideoCaptionMergedDataset)
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
from .parquet_dataset_map_style import build_parquet_map_style_dataloader
__all__ = ["build_parquet_map_style_dataloader"]
from fastvideo.v1.dataset.validation_dataset import ValidationDataset
def getdataset(args, start_idx=0) -> T2V_dataset:
def getdataset(args) -> VideoCaptionMergedDataset:
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
resize_topcrop = [
@@ -31,15 +29,14 @@ def getdataset(args, start_idx=0) -> T2V_dataset:
*resize_topcrop,
norm_fun,
])
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
if args.dataset == "t2v":
return T2V_dataset(args,
transform=transform,
temporal_sample=temporal_sample,
tokenizer=tokenizer,
transform_topcrop=transform_topcrop,
start_idx=start_idx)
return VideoCaptionMergedDataset(data_merge_path=args.data_merge_path,
args=args,
transform=transform,
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop)
raise NotImplementedError(args.dataset)
__all__ = [
"build_parquet_map_style_dataloader", "ValidationDataset",
"VideoCaptionMergedDataset"
]
+60 -11
View File
@@ -26,15 +26,47 @@ pyarrow_schema_i2v = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
pa.field("first_frame_latent_bytes", pa.binary()),
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
pa.field("first_frame_latent_dtype", pa.string()),
# I2V Validation
pa.field("pil_image_bytes", pa.binary()),
pa.field("pil_image_shape", pa.list_(pa.int64())),
pa.field("pil_image_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_i2v_validation = 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()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
# I2V Validation
pa.field("pil_image_bytes", pa.binary()),
pa.field("pil_image_shape", pa.list_(pa.int64())),
pa.field("pil_image_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
@@ -64,11 +96,6 @@ pyarrow_schema_t2v = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
@@ -80,4 +107,26 @@ pyarrow_schema_t2v = pa.schema([
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
])
pyarrow_schema_t2v_validation = 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("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
@@ -4,6 +4,7 @@ import random
from typing import Dict, List, Tuple
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
import tqdm
@@ -70,10 +71,12 @@ class LatentsParquetIterStyleDataset(IterableDataset):
drop_last: bool = True,
text_padding_length: int = 512,
seed: int = 42,
read_batch_size: int = 32):
read_batch_size: int = 32,
parquet_schema: pa.Schema = None):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.parquet_schema = parquet_schema
self.cfg_rate = cfg_rate
self.text_padding_length = text_padding_length
self.seed = seed
@@ -3,6 +3,7 @@ import os
import pickle
from typing import Any, Dict, List, Tuple
import pyarrow as pa
import pyarrow.parquet as pq
# Torch in general
import torch
@@ -11,7 +12,7 @@ import tqdm
from torch.utils.data import Dataset, Sampler
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
from fastvideo.v1.dataset.utils import collate_rows_from_parquet_schema
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
get_world_size)
from fastvideo.v1.logger import init_logger
@@ -185,12 +186,14 @@ class LatentsParquetMapStyleDataset(Dataset):
Using parquet for map style dataset is not efficient, we mainly keep it for backward compatibility and debugging.
"""
# Modify this in the future if we want to add more keys, for example, in image to video.
keys = [("vae_latent", "latent"), "text_embedding"]
keys = [("vae_latent", "latent"), "text_embedding", "clip_feature",
"first_frame_latent", "pil_image"]
def __init__(
self,
path: str,
batch_size: int,
parquet_schema: pa.Schema,
cfg_rate: float = 0.0,
seed: int = 42,
drop_last: bool = True,
@@ -200,6 +203,7 @@ class LatentsParquetMapStyleDataset(Dataset):
super().__init__()
self.path = path
self.cfg_rate = cfg_rate
self.parquet_schema = parquet_schema
if cfg_rate > 0.0:
raise ValueError(
"cfg_rate > 0.0 is not supported for now because it will trigger bug when num_data_workers > 0"
@@ -209,15 +213,6 @@ class LatentsParquetMapStyleDataset(Dataset):
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
self.batch = batch_size
self.text_padding_length = text_padding_length
self._cols = [
"vae_latent_bytes",
"vae_latent_shape",
"text_embedding_bytes",
"text_embedding_shape",
"text_embedding_dtype",
"height",
"width",
]
self.sampler = DP_SP_BatchSampler(
batch_size=batch_size,
dataset_size=sum(self.lengths),
@@ -232,7 +227,7 @@ class LatentsParquetMapStyleDataset(Dataset):
len(self.parquet_files), sum(self.lengths))
def get_validation_negative_prompt(
self) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, str]:
self) -> tuple[torch.Tensor, torch.Tensor, str]:
"""
Get the negative prompt for validation.
This method ensures the negative prompt is loaded and cached properly.
@@ -246,19 +241,22 @@ class LatentsParquetMapStyleDataset(Dataset):
row_dict = read_row_from_parquet_file([file_path], row_idx,
[self.lengths[0]])
all_latents_list, all_embs_list, all_masks_list, caption_text_list = collate_latents_embs_masks(
[row_dict], self.text_padding_length, self.keys)
all_latents, all_embs, all_masks, caption_text = all_latents_list[
0], all_embs_list[0], all_masks_list[0], caption_text_list[0]
# add batch dimension
if len(all_embs.shape) == 2:
all_embs = all_embs.unsqueeze(0)
if len(all_masks.shape) == 1:
all_masks = all_masks.unsqueeze(0).unsqueeze(0)
return all_latents, all_embs, all_masks, caption_text
batch = collate_rows_from_parquet_schema([row_dict],
self.parquet_schema,
self.text_padding_length)
negative_prompt = batch['info_list'][0]['prompt']
negative_prompt_embedding = batch['text_embedding']
negative_prompt_attention_mask = batch['text_attention_mask']
if len(negative_prompt_embedding.shape) == 2:
negative_prompt_embedding = negative_prompt_embedding.unsqueeze(0)
if len(negative_prompt_attention_mask.shape) == 1:
negative_prompt_attention_mask = negative_prompt_attention_mask.unsqueeze(
0).unsqueeze(0)
return negative_prompt_embedding, negative_prompt_attention_mask, negative_prompt
# PyTorch calls this ONLY because the batch_sampler yields a list
def __getitems__(self, indices: List[int]):
def __getitems__(self, indices: List[int]) -> Dict[str, Any]:
"""
Batch fetch using read_row_from_parquet_file for each index.
"""
@@ -267,9 +265,12 @@ class LatentsParquetMapStyleDataset(Dataset):
for idx in indices
]
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
rows, self.text_padding_length, self.keys)
return all_latents, all_embs, all_masks, caption_text
# all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos = collate_latents_embs_masks(
# rows, self.text_padding_length, self.keys)
# return all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos
batch = collate_rows_from_parquet_schema(rows, self.parquet_schema,
self.text_padding_length)
return batch
def __len__(self):
return sum(self.lengths)
@@ -286,6 +287,7 @@ def build_parquet_map_style_dataloader(
path,
batch_size,
num_data_workers,
parquet_schema,
cfg_rate=0.0,
drop_last=True,
drop_first_row=False,
@@ -298,6 +300,7 @@ def build_parquet_map_style_dataloader(
drop_last=drop_last,
drop_first_row=drop_first_row,
text_padding_length=text_padding_length,
parquet_schema=parquet_schema,
seed=seed)
loader = StatefulDataLoader(
@@ -0,0 +1,592 @@
# SPDX-License-Identifier: Apache-2.0
import json
import math
import os
import random
from abc import ABC, abstractmethod
from collections import Counter
from dataclasses import dataclass
from os.path import join as opj
from typing import Any, Dict, List, Optional, Union
import numpy as np
import torch
import torchvision
from einops import rearrange
from PIL import Image
from transformers import AutoTokenizer
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@dataclass
class PreprocessBatch:
"""
Batch information for dataset processing stages.
This class holds all the information about a video-caption or image-caption pair
as it moves through the processing pipeline. Fields are populated by different stages.
"""
# Raw metadata
path: str
cap: Union[str, List[str]]
resolution: Optional[Dict] = None
fps: Optional[float] = None
duration: Optional[float] = None
# Processed metadata
num_frames: Optional[int] = None
sample_frame_index: Optional[List[int]] = None
sample_num_frames: Optional[int] = None
# Processed data
pixel_values: Optional[torch.Tensor] = None
text: Optional[str] = None
input_ids: Optional[torch.Tensor] = None
cond_mask: Optional[torch.Tensor] = None
@property
def is_video(self) -> bool:
"""Check if this is a video item."""
return self.path.endswith(".mp4")
@property
def is_image(self) -> bool:
"""Check if this is an image item."""
return self.path.endswith(".jpg")
class DatasetStage(ABC):
"""
Abstract base class for dataset processing stages.
Similar to PipelineStage but designed for dataset preprocessing operations.
"""
@abstractmethod
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process the dataset batch.
Args:
batch: Dataset batch to process
**kwargs: Additional processing parameters
Returns:
Processed batch
"""
raise NotImplementedError
class DatasetFilterStage(ABC):
"""
Abstract base class for dataset filtering stages.
These stages can filter out items during metadata processing.
"""
@abstractmethod
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if batch should be kept.
Args:
batch: Dataset batch to check
**kwargs: Additional parameters
Returns:
True if batch should be kept, False otherwise
"""
raise NotImplementedError
@abstractmethod
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process the dataset batch (for non-filtering operations).
Args:
batch: Dataset batch to process
**kwargs: Additional processing parameters
Returns:
Processed batch
"""
raise NotImplementedError
class DataValidationStage(DatasetFilterStage):
"""Stage for validating data items."""
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Validate data item.
Args:
batch: Dataset batch to validate
Returns:
True if valid, False if invalid
"""
# Check for caption
if batch.cap is None:
return False
if batch.is_video:
# Validate video-specific fields
if batch.duration is None or batch.fps is None:
return False
elif not batch.is_image:
return False
return True
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""Process does nothing for validation - filtering is handled by should_keep."""
return batch
class ResolutionFilterStage(DatasetFilterStage):
"""Stage for filtering data items based on resolution constraints."""
def __init__(self,
max_h_div_w_ratio: float = 17 / 16,
min_h_div_w_ratio: float = 8 / 16,
max_height: int = 1024,
max_width: int = 1024):
self.max_h_div_w_ratio = max_h_div_w_ratio
self.min_h_div_w_ratio = min_h_div_w_ratio
self.max_height = max_height
self.max_width = max_width
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if data item passes resolution filtering.
Args:
batch: Dataset batch with resolution information
Returns:
True if passes filter, False otherwise
"""
# Only apply to videos
if not batch.is_video:
return True
if batch.resolution is None:
return False
height = batch.resolution.get("height", None)
width = batch.resolution.get("width", None)
if height is None or width is None:
return False
# Check aspect ratio
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
return self.filter_resolution(
height,
width,
max_h_div_w_ratio=hw_aspect_thr * aspect,
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""Process does nothing for resolution filtering - filtering is handled by should_keep."""
return batch
def filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
min_h_div_w_ratio: float) -> bool:
"""Filter based on height/width ratio."""
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
class FrameSamplingStage(DatasetFilterStage):
"""Stage for temporal frame sampling and indexing."""
def __init__(self,
num_frames: int,
train_fps: int,
speed_factor: int = 1,
video_length_tolerance_range: float = 5.0,
drop_short_ratio: float = 0.0):
self.num_frames = num_frames
self.train_fps = train_fps
self.speed_factor = speed_factor
self.video_length_tolerance_range = video_length_tolerance_range
self.drop_short_ratio = drop_short_ratio
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if video should be kept based on length constraints.
Args:
batch: Dataset batch
Returns:
True if should be kept, False otherwise
"""
if batch.is_image:
return True
if batch.duration is None or batch.fps is None:
return False
num_frames = math.ceil(batch.fps * batch.duration)
# Check if video is too long
if (num_frames / batch.fps > self.video_length_tolerance_range *
(self.num_frames / self.train_fps * self.speed_factor)):
return False
# Resample frame indices to check length
frame_interval = batch.fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, num_frames,
frame_interval).astype(int)
# Filter short videos
return not (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio)
def process(self,
batch: PreprocessBatch,
temporal_sample_fn=None,
**kwargs) -> PreprocessBatch:
"""
Process frame sampling for video data items.
Args:
batch: Dataset batch
temporal_sample_fn: Function for temporal sampling
Returns:
Updated batch with frame sampling info
"""
if batch.is_image:
# For images, just add sample info
batch.sample_frame_index = [0]
batch.sample_num_frames = 1
return batch
assert batch.duration is not None and batch.fps is not None
batch.num_frames = math.ceil(batch.fps * batch.duration)
# Resample frame indices
frame_interval = batch.fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, batch.num_frames,
frame_interval).astype(int)
# Temporal crop if too long
if len(frame_indices
) > self.num_frames and temporal_sample_fn is not None:
begin_index, end_index = temporal_sample_fn(len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
batch.sample_frame_index = frame_indices.tolist()
batch.sample_num_frames = len(frame_indices)
return batch
class VideoTransformStage(DatasetStage):
"""Stage for video data transformation."""
def __init__(self, transform) -> None:
self.transform = transform
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Transform video data.
Args:
batch: Dataset batch with video information
Returns:
Batch with transformed video tensor
"""
if not batch.is_video:
return batch
assert os.path.exists(batch.path), f"file {batch.path} do not exist!"
assert batch.sample_frame_index is not None, "Frame indices must be set before transformation"
torchvision_video, _, metadata = torchvision.io.read_video(
batch.path, output_format="TCHW")
video = torchvision_video[batch.sample_frame_index]
if self.transform is not None:
video = self.transform(video)
video = rearrange(video, "t c h w -> c t h w")
video = video.to(torch.uint8)
h, w = video.shape[-2:]
assert (
h / w <= 17 / 16 and h / w >= 8 / 16
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({batch.path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
video = video.float() / 127.5 - 1.0
batch.pixel_values = video
return batch
class ImageTransformStage(DatasetStage):
"""Stage for image data transformation."""
def __init__(self, transform, transform_topcrop) -> None:
self.transform = transform
self.transform_topcrop = transform_topcrop
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Transform image data.
Args:
batch: Dataset batch with image information
Returns:
Batch with transformed image tensor
"""
if not batch.is_image:
return batch
image = Image.open(batch.path).convert("RGB")
image = torch.from_numpy(np.array(image))
image = rearrange(image, "h w c -> c h w").unsqueeze(0)
if self.transform_topcrop is not None:
image = self.transform_topcrop(image)
elif self.transform is not None:
image = self.transform(image)
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
batch.pixel_values = image
return batch
class TextEncodingStage(DatasetStage):
"""Stage for text tokenization and encoding."""
def __init__(self, tokenizer, text_max_length: int, cfg_rate: float = 0.0):
self.tokenizer = tokenizer
self.text_max_length = text_max_length
self.cfg_rate = cfg_rate
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process text data.
Args:
batch: Dataset batch with caption information
Returns:
Batch with encoded text information
"""
text = batch.cap
if not isinstance(text, list):
text = [text]
text = [random.choice(text)]
text = text[0] if random.random() > self.cfg_rate else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
batch.text = text
batch.input_ids = text_tokens_and_mask["input_ids"]
batch.cond_mask = text_tokens_and_mask["attention_mask"]
return batch
class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
torch.distributed.checkpoint.stateful.Stateful):
"""
Merged dataset for video and caption data with stage-based processing.
This dataset processes video and image data through a series of stages:
- Data validation
- Resolution filtering
- Frame sampling
- Transformation
- Text encoding
"""
def __init__(self,
data_merge_path: str,
args,
transform,
temporal_sample,
transform_topcrop,
start_idx: int = 0):
self.data_merge_path = data_merge_path
self.start_idx = start_idx
self.args = args
self.temporal_sample = temporal_sample
# Initialize tokenizer
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
# Initialize processing stages
self._init_stages(args, transform, transform_topcrop, tokenizer)
# Process metadata
self.processed_batches = self._process_metadata()
def _init_stages(self, args, transform, transform_topcrop,
tokenizer) -> None:
"""Initialize all processing stages."""
self.validation_stage = DataValidationStage()
self.resolution_filter_stage = ResolutionFilterStage(
max_height=args.max_height, max_width=args.max_width)
self.frame_sampling_stage = FrameSamplingStage(
num_frames=args.num_frames,
train_fps=args.train_fps,
speed_factor=args.speed_factor,
video_length_tolerance_range=args.video_length_tolerance_range,
drop_short_ratio=args.drop_short_ratio)
self.video_transform_stage = VideoTransformStage(transform)
self.image_transform_stage = ImageTransformStage(
transform, transform_topcrop)
self.text_encoding_stage = TextEncodingStage(
tokenizer=tokenizer,
text_max_length=args.text_max_length,
cfg_rate=args.cfg)
def _load_raw_data(self) -> List[Dict]:
"""Load raw data from JSON files."""
all_data = []
# Read folder-annotation pairs
with open(self.data_merge_path) as f:
folder_anno_pairs = [
line.strip().split(",") for line in f if line.strip()
]
# Process each folder-annotation pair
for folder, annotation_file in folder_anno_pairs:
with open(annotation_file) as f:
data_items = json.load(f)
# Update paths with folder prefix
for item in data_items:
item["path"] = opj(folder, item["path"])
all_data.extend(data_items)
return all_data[self.start_idx:]
def _process_metadata(self) -> List[PreprocessBatch]:
"""Process the raw metadata through all filtering stages."""
raw_data = self._load_raw_data()
processed_batches = []
# Initialize counters
filter_counts = {
"validation_failed": 0,
"resolution_failed": 0,
"frame_sampling_failed": 0
}
sample_num_frames: List[int] = []
for item in raw_data:
batch = PreprocessBatch(path=item["path"],
cap=item["cap"],
resolution=item.get("resolution"),
fps=item.get("fps"),
duration=item.get("duration"))
# Apply filtering stages
if not self._apply_filter_stages(batch, filter_counts):
continue
# Apply frame sampling processing
batch = self.frame_sampling_stage.process(
batch, temporal_sample_fn=self.temporal_sample)
processed_batches.append(batch)
assert batch.sample_num_frames is not None
sample_num_frames.append(batch.sample_num_frames)
self._log_filtering_stats(filter_counts, sample_num_frames,
len(raw_data), len(processed_batches))
return processed_batches
def _apply_filter_stages(self, batch: PreprocessBatch,
filter_counts: Dict[str, int]) -> bool:
"""Apply all filter stages and update counters. Returns True if batch should be kept."""
if not self.validation_stage.should_keep(batch):
filter_counts["validation_failed"] += 1
return False
if not self.resolution_filter_stage.should_keep(batch):
filter_counts["resolution_failed"] += 1
return False
if not self.frame_sampling_stage.should_keep(batch):
filter_counts["frame_sampling_failed"] += 1
return False
return True
def _log_filtering_stats(self, filter_counts: Dict[str, int],
sample_num_frames: List[int], before_count: int,
after_count: int):
"""Log filtering statistics."""
logger.info(
"validation_failed: %d, resolution_failed: %d, frame_sampling_failed: %d, "
"Counter(sample_num_frames): %s, before filter: %d, after filter: %d",
filter_counts['validation_failed'],
filter_counts['resolution_failed'],
filter_counts['frame_sampling_failed'], Counter(sample_num_frames),
before_count, after_count)
def __iter__(self):
"""Iterate through processed data items."""
for idx in range(len(self.processed_batches)):
yield self._get_item(idx)
def __len__(self):
return len(self.processed_batches)
def _get_item(self, idx: int) -> Dict:
"""Get a single processed data item."""
batch = self.processed_batches[idx]
# Apply transformation stages
batch = self.video_transform_stage.process(batch)
batch = self.image_transform_stage.process(batch)
batch = self.text_encoding_stage.process(batch)
# Build result dictionary
result = {
"pixel_values": batch.pixel_values,
"text": batch.text,
"input_ids": batch.input_ids,
"cond_mask": batch.cond_mask,
"path": batch.path,
}
# Add video-specific fields
if batch.is_video:
result.update({"fps": batch.fps, "duration": batch.duration})
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"]
-352
View File
@@ -1,352 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import json
import math
import os
import random
from collections import Counter
from os.path import join as opj
import numpy as np
import torch
import torchvision
from einops import rearrange
from PIL import Image
from torch.utils.data import Dataset
from fastvideo.utils.dataset_utils import DecordInit
from fastvideo.utils.logging_ import main_print
class SingletonMeta(type):
_instances: dict[type, 'SingletonMeta'] = {}
def __call__(cls, *args, **kwargs):
if cls not in cls._instances:
instance = super().__call__(*args, **kwargs)
cls._instances[cls] = instance
return cls._instances[cls]
class DataSetProg(metaclass=SingletonMeta):
def __init__(self) -> None:
self.cap_list: list[dict] = []
self.elements: list[int] = []
self.num_workers = 1
self.n_elements = 0
self.worker_elements: dict[int, list[int]] = {}
self.n_used_elements: dict[int, int] = {}
def set_cap_list(self, num_workers, cap_list, n_elements) -> None:
self.num_workers = num_workers
self.cap_list = cap_list
self.n_elements = n_elements
self.elements = list(range(n_elements))
random.shuffle(self.elements)
print(f"n_elements: {len(self.elements)}", flush=True)
for i in range(self.num_workers):
self.n_used_elements[i] = 0
per_worker = int(
math.ceil(len(self.elements) / float(self.num_workers)))
start = i * per_worker
end = min(start + per_worker, len(self.elements))
self.worker_elements[i] = self.elements[start:end]
def get_item(self, work_info) -> int:
worker_id = 0 if work_info is None else work_info.id
idx = self.worker_elements[worker_id][
self.n_used_elements[worker_id] %
len(self.worker_elements[worker_id])]
self.n_used_elements[worker_id] += 1
return idx
dataset_prog = DataSetProg()
def filter_resolution(h: int,
w: int,
max_h_div_w_ratio: float = 17 / 16,
min_h_div_w_ratio: float = 8 / 16) -> bool:
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
class T2V_dataset(Dataset):
def __init__(self,
args,
transform,
temporal_sample,
tokenizer,
transform_topcrop,
start_idx=0) -> None:
self.start_idx = start_idx
self.data = args.data_merge_path
self.num_frames = args.num_frames
self.train_fps = args.train_fps
self.use_image_num = args.use_image_num
self.transform = transform
self.transform_topcrop = transform_topcrop
self.temporal_sample = temporal_sample
self.tokenizer = tokenizer
self.text_max_length = args.text_max_length
self.cfg = args.cfg
self.speed_factor = args.speed_factor
self.max_height = args.max_height
self.max_width = args.max_width
self.drop_short_ratio = args.drop_short_ratio
assert self.speed_factor >= 1
self.v_decoder = DecordInit()
self.video_length_tolerance_range = args.video_length_tolerance_range
self.support_Chinese = True
if "mt5" not in args.text_encoder_name:
self.support_Chinese = False
cap_list = self.get_cap_list()
assert len(cap_list) > 0
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
self.lengths = self.sample_num_frames
n_elements = len(cap_list)
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
n_elements)
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
def set_checkpoint(self, n_used_elements):
for i in range(len(dataset_prog.n_used_elements)):
dataset_prog.n_used_elements[i] = n_used_elements
def __len__(self):
return dataset_prog.n_elements
def __getitem__(self, idx):
data = self.get_data(idx)
return data
def get_data(self, idx) -> dict:
path = dataset_prog.cap_list[idx]["path"]
if path.endswith(".mp4"):
return self.get_video(idx)
else:
return self.get_image(idx)
def get_video(self, idx) -> dict:
video_path = dataset_prog.cap_list[idx]["path"]
assert os.path.exists(video_path), f"file {video_path} do not exist!"
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
torchvision_video, _, metadata = torchvision.io.read_video(
video_path, output_format="TCHW")
video = torchvision_video[frame_indices]
video = self.transform(video)
video = rearrange(video, "t c h w -> c t h w")
video = video.to(torch.uint8)
assert video.dtype == torch.uint8
h, w = video.shape[-2:]
assert (
h / w <= 17 / 16 and h / w >= 8 / 16
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
video = video.float() / 127.5 - 1.0
text = dataset_prog.cap_list[idx]["cap"]
if not isinstance(text, list):
text = [text]
text = [random.choice(text)]
text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"]
cond_mask = text_tokens_and_mask["attention_mask"]
return dict(pixel_values=video,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=video_path,
fps=dataset_prog.cap_list[idx]["fps"],
duration=dataset_prog.cap_list[idx]["duration"])
def get_image(self, idx) -> dict:
image_data = dataset_prog.cap_list[
idx] # [{'path': path, 'cap': cap}, ...]
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
image = torch.from_numpy(np.array(image)) # [h, w, c]
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
# for i in image:
# h, w = i.shape[-2:]
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
image = (self.transform_topcrop(image) if "human_images"
in image_data["path"] else self.transform(image)
) # [1 C H W] -> num_img [1 C H W]
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
caps: list[str] = (image_data["cap"] if isinstance(
image_data["cap"], list) else [image_data["cap"]])
caps = [random.choice(caps)]
text = caps
input_ids, cond_mask = [], []
single_text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
single_text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"] # 1, l
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
return dict(
pixel_values=image,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=image_data["path"],
)
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
new_cap_list = []
sample_num_frames = []
cnt_too_long = 0
cnt_too_short = 0
cnt_no_cap = 0
cnt_no_resolution = 0
cnt_resolution_mismatch = 0
cnt_movie = 0
cnt_img = 0
for i in cap_list:
path = i["path"]
cap = i.get("cap", None)
# ======no caption=====
if cap is None:
cnt_no_cap += 1
continue
if path.endswith(".mp4"):
# ======no fps and duration=====
duration = i.get("duration", None)
fps = i.get("fps", None)
if fps is None or duration is None:
continue
# ======resolution mismatch=====
resolution = i.get("resolution", None)
if resolution is None:
cnt_no_resolution += 1
continue
else:
if (resolution.get("height", None) is None
or resolution.get("width", None) is None):
cnt_no_resolution += 1
continue
height, width = i["resolution"]["height"], i["resolution"][
"width"]
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
is_pick = filter_resolution(
height,
width,
max_h_div_w_ratio=hw_aspect_thr * aspect,
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
if not is_pick:
print("resolution mismatch")
cnt_resolution_mismatch += 1
continue
# if path == 'finetrainers/3dgs-dissolve/videos/1.mp4':
# from IPython import embed; embed()
i["num_frames"] = math.ceil(fps * duration)
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
if i["num_frames"] / fps > self.video_length_tolerance_range * (
self.num_frames / self.train_fps * self.speed_factor
): # too long video is not suitable for this training stage (self.num_frames)
cnt_too_long += 1
continue
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
frame_interval = fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, i["num_frames"],
frame_interval).astype(int)
# comment out it to enable dynamic frames training
if (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio):
cnt_too_short += 1
continue
# too long video will be temporal-crop randomly
if len(frame_indices) > self.num_frames:
begin_index, end_index = self.temporal_sample(
len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
# frame_indices = frame_indices[:self.num_frames] # head crop
i["sample_frame_index"] = frame_indices.tolist()
new_cap_list.append(i)
i["sample_num_frames"] = len(
i["sample_frame_index"]
) # will use in dataloader(group sampler)
sample_num_frames.append(i["sample_num_frames"])
elif path.endswith(".jpg"): # image
cnt_img += 1
new_cap_list.append(i)
i["sample_num_frames"] = 1
sample_num_frames.append(i["sample_num_frames"])
else:
raise NameError(
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
)
# import ipdb;ipdb.set_trace()
main_print(
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
)
return new_cap_list, sample_num_frames
def decord_read(self, path, frame_indices) -> torch.Tensor:
decord_vr = self.v_decoder(path)
video_data = decord_vr.get_batch(frame_indices).asnumpy()
video_data = torch.from_numpy(video_data)
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
return video_data
def read_jsons(self, data) -> list[dict]:
cap_lists = []
with open(data) as f:
folder_anno = [
i.strip().split(",") for i in f.readlines()
if len(i.strip()) > 0
]
print(folder_anno)
for folder, anno in folder_anno:
with open(anno) as f:
sub_list = json.load(f)
for i in range(len(sub_list)):
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
cap_lists += sub_list
return cap_lists
def get_cap_list(self) -> list:
cap_lists = self.read_jsons(self.data)[self.start_idx:]
return cap_lists
+174 -12
View File
@@ -3,6 +3,10 @@ from typing import Any, Dict, List
import numpy as np
import torch
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
@@ -38,39 +42,70 @@ def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
if shape is None or bytes is None:
raise ValueError(f"Key {key} not found in row_dict")
else:
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
try:
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
except KeyError:
continue
# TODO (peiyuan): read precision
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
if len(data.shape) == 3:
B, L, D = data.shape
assert B == 1, "Batch size must be 1"
data = data.squeeze(0)
return_dict[key] = data
if len(bytes) == 0:
return_dict[key] = torch.zeros(0, dtype=torch.bfloat16)
else:
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
if len(data.shape) == 3:
B, L, D = data.shape
assert B == 1, "Batch size must be 1"
data = data.squeeze(0)
return_dict[key] = data
return return_dict
def collate_latents_embs_masks(
batch_to_process, text_padding_length,
keys) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str]]:
batch_to_process, text_padding_length, keys
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str], Dict[str, Any],
List[Dict[str, Any]]]:
# Initialize tensors to hold padded embeddings and masks
all_latents = []
all_embs = []
all_masks = []
all_clip_features = []
all_first_frame_latents = []
all_pil_images = []
all_infos = []
caption_text = []
# Process each row individually
for i, row in enumerate(batch_to_process):
# Get info from row
info_keys = [
"caption", "file_name", "media_type", "width", "height",
"num_frames", "duration_sec", "fps"
]
info = {}
for key in info_keys:
if key in row:
info[key] = row[key]
else:
info[key] = ""
info["prompt"] = info["caption"]
# Get tensors from row
data = get_torch_tensors_from_row_dict(row, keys)
latents, emb = data["vae_latent"], data["text_embedding"]
clip_feature = data.get("clip_feature", None)
first_frame_latent = data.get("first_frame_latent", None)
pil_image = data.get("pil_image", None)
padded_emb, mask = pad(emb, text_padding_length)
# Store in batch tensors
all_latents.append(latents)
all_embs.append(padded_emb)
all_masks.append(mask)
all_clip_features.append(clip_feature)
all_first_frame_latents.append(first_frame_latent)
all_pil_images.append(pil_image)
all_infos.append(info)
# TODO(py): remove this once we fix preprocess
try:
caption_text.append(row["prompt"])
@@ -81,5 +116,132 @@ def collate_latents_embs_masks(
all_latents = torch.stack(all_latents)
all_embs = torch.stack(all_embs)
all_masks = torch.stack(all_masks)
all_extra_latents = {
"clip_feature": torch.stack(all_clip_features),
"first_frame_latent": torch.stack(all_first_frame_latents),
"pil_image": all_pil_images,
}
return all_latents, all_embs, all_masks, caption_text
return all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos
def collate_rows_from_parquet_schema(rows, parquet_schema,
text_padding_length) -> Dict[str, Any]:
"""
Collate rows from parquet files based on the provided schema.
Dynamically processes tensor fields based on schema and returns batched data.
Args:
rows: List of row dictionaries from parquet files
parquet_schema: PyArrow schema defining the structure of the data
Returns:
Dict containing batched tensors and metadata
"""
if not rows:
return {}
# Initialize containers for different data types
batch_data = {}
# Get tensor and metadata field names from schema (fields ending with '_bytes')
tensor_fields = []
metadata_fields = []
for field in parquet_schema.names:
if field.endswith('_bytes'):
shape_field = field.replace('_bytes', '_shape')
dtype_field = field.replace('_bytes', '_dtype')
tensor_name = field.replace('_bytes', '')
tensor_fields.append(tensor_name)
assert shape_field in parquet_schema.names, f"Shape field {shape_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
assert dtype_field in parquet_schema.names, f"Dtype field {dtype_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
elif not field.endswith('_shape') and not field.endswith('_dtype'):
# Only add actual metadata fields, not the shape/dtype helper fields
metadata_fields.append(field)
# Process each tensor field efficiently
for tensor_name in tensor_fields:
tensor_list = []
for row in rows:
# Get tensor data from row using the existing helper function pattern
shape_key = f"{tensor_name}_shape"
bytes_key = f"{tensor_name}_bytes"
if shape_key in row and bytes_key in row:
# logger.info("row: %s", row)
# logger.info("shape_key: %s", shape_key)
# logger.info("bytes_key: %s", bytes_key)
shape = row[shape_key]
bytes_data = row[bytes_key]
if len(bytes_data) == 0:
tensor = torch.zeros(0, dtype=torch.bfloat16)
else:
# Convert bytes to tensor using float32 as default
# logger.info("len(bytes_data): %s", len(bytes_data))
# logger.info("shape: %s", shape)
data = np.frombuffer(
bytes_data, dtype=np.float32).reshape(shape).copy()
tensor = torch.from_numpy(data)
# if len(data.shape) == 3:
# B, L, D = tensor.shape
# assert B == 1, "Batch size must be 1"
# tensor = tensor.squeeze(0)
tensor_list.append(tensor)
else:
# Handle missing tensor data
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
# Stack tensors with special handling for text embeddings
if tensor_list:
if tensor_name == 'text_embedding':
# Handle text embeddings with padding
padded_tensors = []
attention_masks = []
for tensor in tensor_list:
if tensor.numel() > 0:
padded_tensor, mask = pad(tensor, text_padding_length)
padded_tensors.append(padded_tensor)
attention_masks.append(mask)
else:
# Handle empty embeddings - assume default embedding dimension
padded_tensors.append(
torch.zeros(text_padding_length,
768,
dtype=torch.bfloat16))
attention_masks.append(torch.zeros(text_padding_length))
batch_data[tensor_name] = torch.stack(padded_tensors)
batch_data['text_attention_mask'] = torch.stack(attention_masks)
else:
# Stack other tensors directly, handling None values
valid_tensors = [
t for t in tensor_list if t is not None and t.numel() > 0
]
if valid_tensors:
batch_data[tensor_name] = torch.stack(valid_tensors)
elif tensor_list: # All tensors are empty but exist
batch_data[tensor_name] = torch.stack(tensor_list)
# Process metadata fields efficiently into info_list
info_list = []
for row in rows:
info = {}
for field in metadata_fields:
info[field] = row.get(field, "")
# Add prompt field for backward compatibility
info["prompt"] = info.get("caption", "")
info_list.append(info)
batch_data['info_list'] = info_list
# Add caption_text for backward compatibility
if info_list and 'caption' in info_list[0]:
batch_data['caption_text'] = [info['caption'] for info in info_list]
return batch_data
+104
View File
@@ -0,0 +1,104 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from: https://github.com/a-r-r-o-w/finetrainers/blob/main/finetrainers/data/dataset.py
import os
import pathlib
import datasets
import torch
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vision_utils import load_image, load_video
logger = init_logger(__name__)
class ValidationDataset(torch.utils.data.IterableDataset):
def __init__(self, filename: str):
super().__init__()
self.filename = pathlib.Path(filename)
# get directory of filename
# TODO(will)
self.dir = os.path.abspath(self.filename.parent)
if not self.filename.exists():
raise FileNotFoundError(
f"File {self.filename.as_posix()} does not exist")
if self.filename.suffix == ".csv":
data = datasets.load_dataset("csv",
data_files=self.filename.as_posix(),
split="train")
elif self.filename.suffix == ".json":
data = datasets.load_dataset("json",
data_files=self.filename.as_posix(),
split="train",
field="data")
elif self.filename.suffix == ".parquet":
data = datasets.load_dataset("parquet",
data_files=self.filename.as_posix(),
split="train")
elif self.filename.suffix == ".arrow":
data = datasets.load_dataset("arrow",
data_files=self.filename.as_posix(),
split="train")
else:
_SUPPORTED_FILE_FORMATS = [".csv", ".json", ".parquet", ".arrow"]
raise ValueError(
f"Unsupported file format {self.filename.suffix} for validation dataset. Supported formats are: {_SUPPORTED_FILE_FORMATS}"
)
self._data = data.to_iterable_dataset()
def __iter__(self):
for sample in self._data:
# For consistency reasons, we mandate that "caption" is always present in the validation dataset.
# However, since the model specifications use "prompt", we create an alias here.
sample["prompt"] = sample["caption"]
# Load image or video if the path is provided
# TODO(aryan): need to handle custom columns here for control conditions
sample["image"] = None
sample["video"] = None
if sample.get("image_path", None) is not None:
image_path = sample["image_path"]
image_path = os.path.join(self.dir, image_path)
if not pathlib.Path(image_path).is_file(
) and not image_path.startswith("http"):
logger.warning("Image file %s does not exist.", image_path)
else:
sample["image"] = load_image(image_path)
if sample.get("video_path", None) is not None:
video_path = sample["video_path"]
video_path = os.path.join(self.dir, video_path)
if not pathlib.Path(video_path).is_file(
) and not video_path.startswith("http"):
logger.warning("Video file %s does not exist.", video_path)
else:
sample["video"] = load_video(video_path)
if sample.get("control_image_path", None) is not None:
control_image_path = sample["control_image_path"]
control_image_path = os.path.join(self.dir, control_image_path)
if not pathlib.Path(control_image_path).is_file(
) and not control_image_path.startswith("http"):
logger.warning("Control Image file %s does not exist.",
control_image_path)
else:
sample["control_image"] = load_image(control_image_path)
if sample.get("control_video_path", None) is not None:
control_video_path = sample["control_video_path"]
control_video_path = os.path.join(self.dir, control_video_path)
if not pathlib.Path(control_video_path).is_file(
) and not control_video_path.startswith("http"):
logger.warning("Control Video file %s does not exist.",
control_video_path)
else:
sample["control_video"] = load_video(control_video_path)
sample = {k: v for k, v in sample.items() if v is not None}
yield sample
+24 -11
View File
@@ -8,7 +8,7 @@ from contextlib import contextmanager
from dataclasses import field
from typing import Any, Dict, List, Optional
from fastvideo.v1.configs.pipelines.base import PipelineConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig, STA_Mode
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import FlexibleArgumentParser, StoreBoolean
@@ -63,7 +63,7 @@ class FastVideoArgs:
# STA (Sliding Tile Attention) parameters
mask_strategy_file_path: Optional[str] = None
STA_mode: Optional[str] = None
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
skip_time_steps: int = 15
# Compilation
@@ -74,6 +74,9 @@ class FastVideoArgs:
# VSA parameters
VSA_sparsity: float = 0.0 # inference/validation sparsity
# Stage verification
enable_stage_verification: bool = True
@property
def training_mode(self) -> bool:
return not self.inference_mode
@@ -178,12 +181,10 @@ class FastVideoArgs:
parser.add_argument(
"--STA-mode",
type=str,
default=FastVideoArgs.STA_mode,
choices=[
"STA_inference", "STA_searching", "STA_tuning",
"STA_tuning_cfg", None
],
help="STA mode",
default=FastVideoArgs.STA_mode.value,
choices=[mode.value for mode in STA_Mode],
help=
"STA mode contains STA_inference, STA_searching, STA_tuning, STA_tuning_cfg, None",
)
parser.add_argument(
"--skip-time-steps",
@@ -231,6 +232,14 @@ class FastVideoArgs:
help="Validation sparsity for VSA",
)
# Stage verification
parser.add_argument(
"--enable-stage-verification",
action=StoreBoolean,
default=FastVideoArgs.enable_stage_verification,
help="Enable input/output verification for pipeline stages",
)
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
@@ -379,7 +388,8 @@ class TrainingArgs(FastVideoArgs):
precondition_outputs: bool = False
# validation & logs
validation_prompt_dir: str = ""
validation_dataset_file: str = ""
validation_preprocessed_path: str = ""
validation_sampling_steps: str = ""
validation_guidance_scale: str = ""
validation_steps: float = 0.0
@@ -527,9 +537,12 @@ class TrainingArgs(FastVideoArgs):
help="Whether to precondition the outputs of the model")
# Validation and logging
parser.add_argument("--validation-prompt-dir",
parser.add_argument("--validation-dataset-file",
type=str,
help="Directory containing validation prompts")
help="Path to unprocessed validation dataset")
parser.add_argument("--validation-preprocessed-path",
type=str,
help="Path to processed validation dataset")
parser.add_argument("--validation-sampling-steps",
type=str,
help="Validation sampling steps")
+9 -8
View File
@@ -5,15 +5,16 @@ import time
from collections import defaultdict
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Optional
from typing import TYPE_CHECKING, Optional
import torch
# if TYPE_CHECKING:
from fastvideo.v1.attention import AttentionMetadata
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
if TYPE_CHECKING:
from fastvideo.v1.attention import AttentionMetadata
from fastvideo.v1.pipelines import ForwardBatch
logger = init_logger(__name__)
@@ -36,13 +37,13 @@ class ForwardContext:
# attn_layers: Dict[str, Any]
# TODO: extend to support per-layer dynamic forward context
attn_metadata: "AttentionMetadata" # set dynamically for each forward pass
forward_batch: Optional[ForwardBatch] = None
forward_batch: Optional["ForwardBatch"] = None
_forward_context: Optional[ForwardContext] = None
_forward_context: Optional["ForwardContext"] = None
def get_forward_context() -> ForwardContext:
def get_forward_context() -> "ForwardContext":
"""Get the current forward context."""
assert _forward_context is not None, (
"Forward context is not set. "
@@ -54,7 +55,7 @@ def get_forward_context() -> ForwardContext:
@contextmanager
def set_forward_context(current_timestep,
attn_metadata,
forward_batch: Optional[ForwardBatch] = None,
forward_batch: Optional["ForwardBatch"] = None,
fastvideo_args: Optional[FastVideoArgs] = None):
"""A context manager that stores the current forward context,
can be attention metadata, etc.
+42 -9
View File
@@ -5,6 +5,7 @@ from typing import Optional, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.v1.layers.custom_op import CustomOp
@@ -95,6 +96,22 @@ class ScaleResidual(nn.Module):
return residual + x * gate
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
# NOTE(will): Needed to match behavior of diffusers and wan2.1 even while using
# FSDP's MixedPrecisionPolicy
class FP32LayerNorm(nn.LayerNorm):
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
origin_dtype = inputs.dtype
return F.layer_norm(
inputs.float(),
self.normalized_shape,
self.weight.float() if self.weight is not None else None,
self.bias.float() if self.bias is not None else None,
self.eps,
).to(origin_dtype)
class ScaleResidualLayerNormScaleShift(nn.Module):
"""
Fused operation that combines:
@@ -112,6 +129,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
eps: float = 1e-6,
elementwise_affine: bool = False,
dtype: torch.dtype = torch.float32,
compute_dtype: torch.dtype | None = None,
prefix: str = "",
):
super().__init__()
@@ -121,10 +139,15 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
eps=eps,
dtype=dtype)
elif norm_type == "layer":
self.norm = nn.LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps,
dtype=dtype)
if compute_dtype == torch.float32:
self.norm = FP32LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps)
else:
self.norm = nn.LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps,
dtype=dtype)
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
@@ -163,18 +186,25 @@ class LayerNormScaleShift(nn.Module):
eps: float = 1e-6,
elementwise_affine: bool = False,
dtype: torch.dtype = torch.float32,
compute_dtype: torch.dtype | None = None,
prefix: str = "",
):
super().__init__()
self.compute_dtype = compute_dtype
if norm_type == "rms":
self.norm = RMSNorm(hidden_size,
has_weight=elementwise_affine,
eps=eps)
elif norm_type == "layer":
self.norm = nn.LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps,
dtype=dtype)
if self.compute_dtype == torch.float32:
self.norm = FP32LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps)
else:
self.norm = nn.LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps,
dtype=dtype)
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
@@ -182,4 +212,7 @@ class LayerNormScaleShift(nn.Module):
scale: torch.Tensor) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
normalized = self.norm(x)
return normalized * (1.0 + scale) + shift
if self.compute_dtype == torch.float32:
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
else:
return normalized * (1.0 + scale) + shift
+31 -12
View File
@@ -14,8 +14,8 @@ from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.configs.sample.wan import WanTeaCacheParams
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
from fastvideo.v1.forward_context import get_forward_context
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
ScaleResidual,
from fastvideo.v1.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.v1.layers.linear import ReplicatedLinear
# from torch.nn import RMSNorm
@@ -25,18 +25,21 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.dits.base import CachableDiT
from fastvideo.v1.platforms import AttentionBackendEnum
logger = init_logger(__name__)
class WanImageEmbedding(torch.nn.Module):
def __init__(self, in_features: int, out_features: int):
super().__init__()
self.norm1 = nn.LayerNorm(in_features)
self.norm1 = FP32LayerNorm(in_features)
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
self.norm2 = nn.LayerNorm(out_features)
self.norm2 = FP32LayerNorm(out_features)
def forward(self,
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
@@ -232,7 +235,7 @@ class WanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -263,7 +266,8 @@ class WanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
@@ -283,7 +287,8 @@ class WanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -375,7 +380,7 @@ class WanTransformerBlock_VSA(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -407,7 +412,8 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
@@ -427,7 +433,8 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -564,7 +571,8 @@ class WanTransformer3DModel(CachableDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -572,6 +580,17 @@ class WanTransformer3DModel(CachableDiT):
self.gradient_checkpointing = False
# For type checking
self.previous_e0_even = None
self.previous_e0_odd = None
self.previous_residual_even = None
self.previous_residual_odd = None
self.is_even = True
self.should_calc_even = True
self.should_calc_odd = True
self.accumulated_rel_l1_distance_even = 0
self.accumulated_rel_l1_distance_odd = 0
self.cnt = 0
self.__post_init__()
def forward(self,
@@ -658,7 +677,7 @@ class WanTransformer3DModel(CachableDiT):
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
-5
View File
@@ -51,11 +51,6 @@ def auto_attributes(init_func):
return wrapper
def set_random_seed(seed: int) -> None:
from fastvideo.v1.platforms import current_platform
current_platform.seed_everything(seed)
def set_weight_attrs(
weight: torch.Tensor,
weight_attrs: Optional[Dict[str, Any]],
+83
View File
@@ -1,8 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
import os
import tempfile
from typing import Callable, List, Optional, Tuple, Union
from urllib.parse import unquote, urlparse
import imageio
import numpy as np
import PIL.Image
import PIL.ImageOps
@@ -86,6 +89,7 @@ def normalize(
return 2.0 * images - 1.0
# adapted from diffusers.utils import load_image
def load_image(
image: Union[str, PIL.Image.Image],
convert_method: Optional[Callable[[PIL.Image.Image],
@@ -131,6 +135,85 @@ def load_image(
return image
# adapted from diffusers.utils import load_video
def load_video(
video: str,
convert_method: Optional[Callable[[List[PIL.Image.Image]],
List[PIL.Image.Image]]] = None,
) -> List[PIL.Image.Image]:
"""
Loads `video` to a list of PIL Image.
Args:
video (`str`):
A URL or Path to a video to convert to a list of PIL Image format.
convert_method (Callable[[List[PIL.Image.Image]], List[PIL.Image.Image]], *optional*):
A conversion method to apply to the video after loading it. When set to `None` the images will be converted
to "RGB".
Returns:
`List[PIL.Image.Image]`:
The video as a list of PIL images.
"""
is_url = video.startswith("http://") or video.startswith("https://")
is_file = os.path.isfile(video)
was_tempfile_created = False
if not (is_url or is_file):
raise ValueError(
f"Incorrect path or URL. URLs must start with `http://` or `https://`, and {video} is not a valid path."
)
if is_url:
response = requests.get(video, stream=True)
if response.status_code != 200:
raise ValueError(
f"Failed to download video. Status code: {response.status_code}"
)
parsed_url = urlparse(video)
file_name = os.path.basename(unquote(parsed_url.path))
suffix = os.path.splitext(file_name)[1] or ".mp4"
with tempfile.NamedTemporaryFile(suffix=suffix,
delete=False) as temp_file:
video_path = temp_file.name
video_data = response.iter_content(chunk_size=8192)
for chunk in video_data:
temp_file.write(chunk)
video = video_path
pil_images = []
if video.endswith(".gif"):
gif = PIL.Image.open(video)
try:
while True:
pil_images.append(gif.copy())
gif.seek(gif.tell() + 1)
except EOFError:
pass
else:
try:
imageio.plugins.ffmpeg.get_exe()
except AttributeError:
raise AttributeError(
"`Unable to find an ffmpeg installation on your machine. Please install via `pip install imageio-ffmpeg"
) from None
with imageio.get_reader(video) as reader:
# Read all frames
for frame in reader:
pil_images.append(PIL.Image.fromarray(frame))
if was_tempfile_created:
os.remove(video_path)
if convert_method is not None:
pil_images = convert_method(pil_images)
return pil_images
def get_default_height_width(
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
vae_scale_factor: int,
+3 -1
View File
@@ -11,7 +11,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
TrainingBatch)
from fastvideo.v1.pipelines.pipeline_registry import PipelineRegistry
from fastvideo.v1.utils import (maybe_download_model,
verify_model_config_and_directory)
@@ -63,4 +64,5 @@ __all__ = [
"PipelineRegistry",
"ForwardBatch",
"LoRAPipeline",
"TrainingBatch",
]
@@ -298,3 +298,7 @@ class ComposedPipelineBase(ABC):
# Return the output
return batch
def train(self) -> None:
raise NotImplementedError(
"if training_mode is True, the pipeline must implement this method")
@@ -11,8 +11,10 @@ import pprint
from dataclasses import asdict, dataclass, field
from typing import Any, Dict, List, Optional, Union
import PIL.Image
import torch
from fastvideo.v1.attention import AttentionMetadata
from fastvideo.v1.configs.sample.teacache import (TeaCacheParams,
WanTeaCacheParams)
@@ -36,6 +38,8 @@ class ForwardBatch:
# Image inputs
image_path: Optional[str] = None
image_embeds: List[torch.Tensor] = field(default_factory=list)
pil_image: Optional[PIL.Image.Image] = None
preprocessed_image: Optional[torch.Tensor] = None
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
@@ -136,3 +140,38 @@ class ForwardBatch:
def __str__(self):
return pprint.pformat(asdict(self), indent=2, width=120)
@dataclass
class TrainingBatch:
current_timestep: int = 0
current_vsa_sparsity: float = 0.0
# Dataloader batch outputs
latents: Optional[torch.Tensor] = None
encoder_hidden_states: Optional[torch.Tensor] = None
encoder_attention_mask: Optional[torch.Tensor] = None
# i2v
# extra_latents: Optional[Dict[str, Any]] = None
preprocessed_image: Optional[torch.Tensor] = None
image_embeds: Optional[torch.Tensor] = None
image_latents: Optional[torch.Tensor] = None
infos: Optional[List[Dict[str, Any]]] = None
# Transformer inputs
noisy_model_input: Optional[torch.Tensor] = None
timesteps: Optional[torch.Tensor] = None
sigmas: Optional[torch.Tensor] = None
noise: Optional[torch.Tensor] = None
attn_metadata: Optional[AttentionMetadata] = None
# input kwargs
input_kwargs: Optional[Dict[str, Any]] = None
# Training loss
loss: torch.Tensor | None = None
# Training outputs
total_loss: float | None = None
grad_norm: float | None = None
@@ -2,19 +2,22 @@
import gc
import multiprocessing
import os
from collections import defaultdict
from concurrent.futures import ProcessPoolExecutor
from itertools import chain
from typing import Any, Dict, List, Optional
import numpy as np
import PIL.Image
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset import getdataset
from fastvideo.v1.dataset import ValidationDataset, getdataset
from fastvideo.v1.dataset.preprocessing_datasets import PreprocessBatch
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -46,7 +49,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
# 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_validation_text(fastvideo_args, args)
self.preprocess_validation(fastvideo_args, args)
self.preprocess_video_and_text(fastvideo_args, args)
def get_extra_features(self, valid_data: Dict[str, Any],
@@ -58,39 +61,206 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"""Get the schema fields for the pipeline type. Override in subclasses."""
raise NotImplementedError
def create_record_for_schema(self,
preprocess_batch: PreprocessBatch,
schema: pa.Schema,
strict: bool = False) -> Dict[str, Any]:
"""Create a record for the Parquet dataset using a generic schema-based approach.
Args:
preprocess_batch: The batch containing the data to extract
schema: PyArrow schema defining the expected fields
strict: If True, raises an exception when required fields are missing or unfilled
Returns:
Dictionary record matching the schema
Raises:
ValueError: If strict=True and required fields are missing or unfilled
"""
record = {}
unfilled_fields = []
for field in schema.names:
field_filled = False
if field.endswith('_bytes'):
# Handle binary tensor data - convert numpy array or tensor to bytes
tensor_name = field.replace('_bytes', '')
tensor_data = getattr(preprocess_batch, tensor_name, None)
if tensor_data is not None:
try:
if hasattr(tensor_data, 'numpy'): # torch tensor
record[field] = tensor_data.cpu().numpy().tobytes()
field_filled = True
elif hasattr(tensor_data, 'tobytes'): # numpy array
record[field] = tensor_data.tobytes()
field_filled = True
else:
raise ValueError(
f"Unsupported tensor type for field {field}: {type(tensor_data)}"
)
except Exception as e:
if strict:
raise ValueError(
f"Failed to convert tensor {tensor_name} to bytes: {e}"
)
record[field] = b'' # Empty bytes for missing data
else:
record[field] = b'' # Empty bytes for missing data
elif field.endswith('_shape'):
# Handle tensor shape info
tensor_name = field.replace('_shape', '')
tensor_data = getattr(preprocess_batch, tensor_name, None)
if tensor_data is not None and hasattr(tensor_data, 'shape'):
record[field] = list(tensor_data.shape)
field_filled = True
else:
record[field] = []
elif field.endswith('_dtype'):
# Handle tensor dtype info
tensor_name = field.replace('_dtype', '')
tensor_data = getattr(preprocess_batch, tensor_name, None)
if tensor_data is not None and hasattr(tensor_data, 'dtype'):
record[field] = str(tensor_data.dtype)
field_filled = True
else:
record[field] = 'unknown'
elif field in ['width', 'height', 'num_frames']:
# Handle integer metadata fields
value = getattr(preprocess_batch, field, None)
if value is not None:
try:
record[field] = int(value)
field_filled = True
except (ValueError, TypeError) as e:
if strict:
raise ValueError(
f"Failed to convert field {field} to int: {e}")
record[field] = 0
else:
record[field] = 0
elif field in ['duration_sec', 'fps']:
# Handle float metadata fields
# Map schema field names to batch attribute names
attr_name = 'duration' if field == 'duration_sec' else field
value = getattr(preprocess_batch, attr_name, None)
if value is not None:
try:
record[field] = float(value)
field_filled = True
except (ValueError, TypeError) as e:
if strict:
raise ValueError(
f"Failed to convert field {field} to float: {e}"
)
record[field] = 0.0
else:
record[field] = 0.0
else:
# Handle string fields (id, file_name, caption, media_type, etc.)
# Map common schema field names to batch attribute names
attr_name = field
if field == 'caption':
attr_name = 'text'
elif field == 'file_name':
attr_name = 'path'
elif field == 'id':
# Generate ID from path if available
path_value = getattr(preprocess_batch, 'path', None)
if path_value:
import os
record[field] = os.path.basename(path_value).split(
'.')[0]
field_filled = True
else:
record[field] = ""
continue
elif field == 'media_type':
# Determine media type from path
path_value = getattr(preprocess_batch, 'path', None)
if path_value:
record[field] = 'video' if path_value.endswith(
'.mp4') else 'image'
field_filled = True
else:
record[field] = ""
continue
value = getattr(preprocess_batch, attr_name, None)
if value is not None:
record[field] = str(value)
field_filled = True
else:
record[field] = ""
# Track unfilled fields
if not field_filled:
unfilled_fields.append(field)
# Handle strict mode
if strict and unfilled_fields:
raise ValueError(
f"Required fields were not filled: {unfilled_fields}")
# Log unfilled fields as warning if not in strict mode
if unfilled_fields:
logger.warning(
f"Some fields were not filled and got default values: {unfilled_fields}"
)
return record
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
text_attention_mask: np.ndarray,
valid_data: Optional[Dict[str, Any]],
# text_attention_mask: np.ndarray,
valid_data: Dict[str, Any],
idx: int,
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Create a record for the Parquet dataset."""
record = {
"id": video_name,
"vae_latent_bytes": vae_latent.tobytes(),
"vae_latent_shape": list(vae_latent.shape),
"vae_latent_dtype": str(vae_latent.dtype),
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape": list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx] if valid_data else "",
"media_type": "video",
"id":
video_name,
"vae_latent_bytes":
vae_latent.tobytes(),
"vae_latent_shape":
list(vae_latent.shape),
"vae_latent_dtype":
str(vae_latent.dtype),
"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":
valid_data["text"][idx] if len(valid_data["text"]) > 0 else "",
"media_type":
"video",
"width":
valid_data["pixel_values"][idx].shape[-2] if valid_data else 0,
valid_data["pixel_values"][idx].shape[-2]
if len(valid_data["pixel_values"]) > 0 else 0,
"height":
valid_data["pixel_values"][idx].shape[-1] if valid_data else 0,
valid_data["pixel_values"][idx].shape[-1]
if len(valid_data["pixel_values"]) > 0 else 0,
"num_frames":
vae_latent.shape[1] if len(vae_latent.shape) > 1 else 0,
"duration_sec":
float(valid_data["duration"][idx]) if valid_data else 0.0,
"fps": float(valid_data["fps"][idx]) if valid_data else 0.0,
float(valid_data["duration"][idx])
if len(valid_data["duration"]) > 0 else 0.0,
"fps":
float(valid_data["fps"][idx])
if len(valid_data["fps"]) > 0 else 0.0,
}
if extra_features:
record.update(extra_features)
@@ -103,7 +273,6 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"combined_parquet_dataset")
os.makedirs(combined_parquet_dir, exist_ok=True)
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
# Get how many samples have already been processed
start_idx = 0
@@ -114,14 +283,10 @@ class BasePreprocessPipeline(ComposedPipelineBase):
start_idx += table.num_rows
# Loading dataset
train_dataset = getdataset(args, start_idx=start_idx)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=False)
train_dataset = getdataset(args)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
@@ -215,8 +380,8 @@ class BasePreprocessPipeline(ComposedPipelineBase):
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
).astype(np.uint8)
# text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
# ).astype(np.uint8)
# Get extra features for this sample if needed
sample_extra_features = {}
@@ -233,7 +398,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
# text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
@@ -285,7 +450,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
num_processed_samples = 0
self.all_tables = []
def preprocess_validation_text(self, fastvideo_args: FastVideoArgs, args):
def preprocess_validation(self, fastvideo_args: FastVideoArgs, args):
"""Process validation text prompts and save them to parquet files.
This base implementation handles the common validation text processing logic.
@@ -296,22 +461,32 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"validation_parquet_dataset")
os.makedirs(validation_parquet_dir, exist_ok=True)
with open(args.validation_prompt_txt, encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
validation_dataset = ValidationDataset(args.validation_dataset_file)
# Prepare batch data for Parquet dataset
batch_data = []
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
if sampling_param.negative_prompt:
prompts = [sampling_param.negative_prompt] + prompts
negative_prompt = {
'caption': sampling_param.negative_prompt,
'image_path': None,
'video_path': None,
}
validation_iterable = chain([negative_prompt], validation_dataset)
else:
negative_prompt = None
validation_iterable = validation_dataset
# Add progress bar for validation text preprocessing
pbar = tqdm(enumerate(prompts),
pbar = tqdm(enumerate(validation_iterable),
desc="Processing validation prompts",
unit="prompt")
for prompt_idx, prompt in pbar:
for idx, sample in pbar:
with torch.inference_mode():
prompt = sample["caption"]
is_negative_prompt = idx == 0
# Text Encoder
batch = ForwardBatch(
data_type="video",
@@ -338,15 +513,44 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"Shape after removing padding - Embeddings: %s, Mask: %s",
text_embedding.shape, text_attention_mask.shape)
extra_features = {}
if not is_negative_prompt:
height = sample["height"]
width = sample["width"]
if "image_path" in sample and "video_path" in sample:
raise ValueError(
"Only one of image_path or video_path should be provided"
)
if "image" in sample:
extra_features = self.preprocess_image(
sample["image"], height, width, fastvideo_args)
if "video" in sample:
extra_features = self.preprocess_video(
sample["video"], height, width, fastvideo_args)
# Get extra features for this sample if needed
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
if isinstance(value, torch.Tensor):
sample_extra_features[key] = value.cpu().numpy()
else:
sample_extra_features[key] = value
valid_data = defaultdict(list)
valid_data["text"] = [prompt]
# Create record for Parquet dataset
record = self.create_record(video_name=file_name,
vae_latent=np.array([],
dtype=np.float32),
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
valid_data=None,
idx=0,
extra_features=None)
record = self.create_record(
video_name=file_name,
vae_latent=np.array([], dtype=np.float32),
text_embedding=text_embedding,
# text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=0,
extra_features=sample_extra_features)
batch_data.append(record)
logger.info("Saved validation sample: %s", file_name)
@@ -420,6 +624,15 @@ class BasePreprocessPipeline(ComposedPipelineBase):
del table
gc.collect() # Force garbage collection
def preprocess_image(self, image: PIL.Image.Image, height: int, width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
return {}
def preprocess_video(self, video: list[PIL.Image.Image], height: int,
width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
return {}
def _flush_tables(self, num_processed_samples: int, args,
combined_parquet_dir: str):
"""Flush collected tables to disk."""
@@ -8,6 +8,7 @@ using the modular pipeline architecture.
from typing import Any, Dict, List, Optional
import numpy as np
import PIL
import torch
from PIL import Image
@@ -15,8 +16,13 @@ from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.models.vision_utils import (get_default_height_width,
normalize, numpy_to_pt,
pil_to_numpy, resize)
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.v1.pipelines.stages import ImageEncodingStage, TextEncodingStage
class PreprocessPipeline_I2V(BasePreprocessPipeline):
@@ -26,18 +32,75 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
"text_encoder", "tokenizer", "vae", "image_encoder", "image_processor"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
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="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
def preprocess_image(self, image: PIL.Image.Image, height: int, width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
assert hasattr(
self,
"image_encoding_stage"), "Image encoding stage must be created"
batch = ForwardBatch(
data_type="video",
pil_image=image,
)
result_batch = self.image_encoding_stage(batch, fastvideo_args)
clip_features = result_batch.image_embeds[0]
# image = self.pil_to_tensor(image)
image = self.preprocess(
image,
vae_scale_factor=self.get_module("vae").spatial_compression_ratio,
height=height,
width=width)
return {
"clip_feature": clip_features[0],
"pil_image": image,
}
def preprocess_video(self, video: list[PIL.Image.Image], height: int,
width: int,
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
return self.preprocess_image(video[0], height, width, fastvideo_args)
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_extra_features(self, valid_data: Dict[str, Any],
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
# TODO(will): move these to cpu at some point
self.get_module("image_encoder").to(get_torch_device())
# self.get_module("image_processor").to(get_torch_device())
self.get_module("vae").to(get_torch_device())
features = {}
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
batch_size, _, num_frames, height, width = valid_data[
"pixel_values"].shape
latent_height = height // self.get_module(
"vae").spatial_compression_ratio
latent_width = width // self.get_module("vae").spatial_compression_ratio
processed_images = []
# Frame has values between -1 and 1
for frame in first_frame:
frame = (frame + 1) * 127.5
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
processed_img = self.get_module("image_processor")(
images=frame_pil, return_tensors="pt")
@@ -53,25 +116,90 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
clip_features = self.get_module("image_encoder")(**image_inputs)
clip_features = clip_features.last_hidden_state
return {"clip_feature": clip_features}
features["clip_feature"] = clip_features
"""Get VAE features from the first frame of each video"""
video_conditions = []
for frame in first_frame:
processed_img = frame.to(device="cpu", dtype=torch.float32)
processed_img = processed_img.unsqueeze(0).permute(0, 3, 1,
2).unsqueeze(2)
# (B, H, W, C) -> (B, C, 1, H, W)
video_condition = torch.cat([
processed_img,
processed_img.new_zeros(processed_img.shape[0],
processed_img.shape[1], num_frames - 1,
height, width)
],
dim=2)
video_condition = video_condition.to(device=get_torch_device(),
dtype=torch.float32)
video_conditions.append(video_condition)
video_conditions = torch.cat(video_conditions, dim=0)
with torch.autocast(device_type="cuda",
dtype=torch.float32,
enabled=True):
encoder_outputs = self.get_module("vae").encode(video_conditions)
latent_condition = encoder_outputs.mean
if (hasattr(self.get_module("vae"), "shift_factor")
and self.get_module("vae").shift_factor is not None):
if isinstance(self.get_module("vae").shift_factor, torch.Tensor):
latent_condition -= self.get_module("vae").shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.get_module("vae").shift_factor
if isinstance(self.get_module("vae").scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor.to(latent_condition.device,
latent_condition.dtype)
else:
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
latent_width)
mask_lat_size[:, :, list(range(1, num_frames))] = 0
first_frame_mask = mask_lat_size[:, :, 0:1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask,
dim=2,
repeats=self.get_module("vae").temporal_compression_ratio)
mask_lat_size = torch.concat(
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
mask_lat_size = mask_lat_size.view(
batch_size, -1,
self.get_module("vae").temporal_compression_ratio, latent_height,
latent_width)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(latent_condition.device)
image_latent = torch.concat([mask_lat_size, latent_condition], dim=1)
features["first_frame_latent"] = image_latent
return features
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
text_attention_mask: np.ndarray,
# text_attention_mask: np.ndarray,
valid_data: Optional[Dict[str, Any]],
idx: int,
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Create a record for the Parquet dataset with CLIP features."""
record = super().create_record(video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
record = super().create_record(
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
# text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "clip_feature" in extra_features:
clip_feature = extra_features["clip_feature"]
@@ -87,7 +215,69 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
"clip_feature_dtype": "",
})
return record # type: ignore
if extra_features and "first_frame_latent" in extra_features:
first_frame_latent = extra_features["first_frame_latent"]
record.update({
"first_frame_latent_bytes":
first_frame_latent.tobytes(),
"first_frame_latent_shape":
list(first_frame_latent.shape),
"first_frame_latent_dtype":
str(first_frame_latent.dtype),
})
else:
record.update({
"first_frame_latent_bytes": b"",
"first_frame_latent_shape": [],
"first_frame_latent_dtype": "",
})
if extra_features and "pil_image" in extra_features:
pil_image = extra_features["pil_image"]
record.update({
"pil_image_bytes": pil_image.tobytes(),
"pil_image_shape": list(pil_image.shape),
"pil_image_dtype": str(pil_image.dtype),
})
else:
record.update({
"pil_image_bytes": b"",
"pil_image_shape": [],
"pil_image_dtype": "",
})
return record
def pil_to_tensor(self, image: PIL.Image.Image) -> torch.Tensor:
image = image
image = np.array(image).astype(np.float32)
image = torch.from_numpy(image)
return image
def preprocess(self,
image: PIL.Image.Image,
vae_scale_factor: int,
height: int,
width: int,
resize_mode: str = "default") -> torch.Tensor:
image = [image]
height, width = get_default_height_width(image[0], vae_scale_factor,
height, width)
image = [
resize(i, height, width, resize_mode=resize_mode) for i in image
]
image = pil_to_numpy(image) # to np
image = numpy_to_pt(image) # to pt
do_normalize = True
if image.min() < 0:
do_normalize = False
if do_normalize:
image = normalize(image)
return image
EntryClass = PreprocessPipeline_I2V
@@ -44,7 +44,7 @@ if __name__ == "__main__":
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument("--validation_dataset_file", type=str)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
@@ -79,7 +79,6 @@ 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("--dataset", default="t2v")
parser.add_argument("--preprocess_task", type=str, default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
+109 -17
View File
@@ -15,10 +15,16 @@ import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
class StageVerificationError(Exception):
"""Exception raised when stage verification fails."""
pass
class PipelineStage(ABC):
"""
Abstract base class for all pipeline stages.
@@ -28,6 +34,70 @@ class PipelineStage(ABC):
for a specific part of the process, such as prompt encoding, latent preparation, etc.
"""
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""
Verify the input for the stage.
Example:
from fastvideo.v1.pipelines.stages.validators import V, VerificationResult
def verify_input(self, batch, fastvideo_args):
result = VerificationResult()
result.add_check("height", batch.height, V.positive_int_divisible(8))
result.add_check("width", batch.width, V.positive_int_divisible(8))
result.add_check("image_latent", batch.image_latent, V.is_tensor)
return result
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
A VerificationResult containing the verification status.
"""
# Default implementation - no verification
return VerificationResult()
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""
Verify the output for the stage.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
A VerificationResult containing the verification status.
"""
# Default implementation - no verification
return VerificationResult()
def _run_verification(self, verification_result: VerificationResult,
stage_name: str, verification_type: str) -> None:
"""
Run verification and raise errors if any checks fail.
Args:
verification_result: Results from verify_input or verify_output
stage_name: Name of the current stage
verification_type: "input" or "output"
"""
if not verification_result.is_valid():
failed_fields = verification_result.get_failed_fields()
if failed_fields:
# Get detailed failure information
detailed_summary = verification_result.get_failure_summary()
failed_fields_str = ", ".join(failed_fields)
error_msg = (
f"{verification_type.capitalize()} verification failed for {stage_name}: "
f"Failed fields: {failed_fields_str}\n"
f"Details: {detailed_summary}")
raise StageVerificationError(error_msg)
@property
def device(self) -> torch.device:
"""Get the device for this stage."""
@@ -48,7 +118,7 @@ class PipelineStage(ABC):
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Execute the stage's processing on the batch with optional logging.
Execute the stage's processing on the batch with optional verification and logging.
Should not be overridden by subclasses.
Args:
@@ -58,34 +128,56 @@ class PipelineStage(ABC):
Returns:
The updated batch information after this stage's processing.
"""
# if envs.ENABLE_STAGE_LOGGING:
stage_name = self.__class__.__name__
# Check if verification is enabled (simple approach for prototype)
enable_verification = getattr(fastvideo_args,
'enable_stage_verification', False)
if enable_verification:
# Pre-execution input verification
try:
input_result = self.verify_input(batch, fastvideo_args)
self._run_verification(input_result, stage_name, "input")
except Exception as e:
logger.error("Input verification failed for %s: %s", stage_name,
str(e))
raise
# Execute the actual stage logic
# envs.ENABLE_STAGE_LOGGING
if False:
self._logger.info("[%s] Starting execution", self._stage_name)
self._logger.info("[%s] Starting execution", stage_name)
start_time = time.perf_counter()
try:
# Call the actual implementation
result = self._call_implementation(batch, fastvideo_args)
result = self.forward(batch, fastvideo_args)
execution_time = time.perf_counter() - start_time
self._logger.info("[%s] Execution completed in %s ms",
self._stage_name, execution_time * 1000)
return result
stage_name, execution_time * 1000)
except Exception as e:
execution_time = time.perf_counter() - start_time
self._logger.error(
"[%s] Error during execution after %s ms: %s",
self._stage_name, execution_time * 1000, e)
self._logger.error("[%s] Traceback: %s", self._stage_name,
"[%s] Error during execution after %s ms: %s", stage_name,
execution_time * 1000, e)
self._logger.error("[%s] Traceback: %s", stage_name,
traceback.format_exc())
# Re-raise the exception
raise
else:
# Just call the implementation directly if logging is disabled
# TODO(will): Also handle backward
return self.forward(batch, fastvideo_args)
# Direct execution (current behavior)
result = self.forward(batch, fastvideo_args)
if enable_verification:
# Post-execution output verification
try:
output_result = self.verify_output(result, fastvideo_args)
self._run_verification(output_result, stage_name, "output")
except Exception as e:
logger.error("Output verification failed for %s: %s",
stage_name, str(e))
raise
return result
@abstractmethod
def forward(
@@ -9,6 +9,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
from fastvideo.v1.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
@@ -69,3 +71,24 @@ class ConditioningStage(PipelineStage):
[batch.negative_attention_mask_2, batch.attention_mask_2])
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify conditioning stage inputs."""
result = VerificationResult()
result.add_check("do_classifier_free_guidance",
batch.do_classifier_free_guidance, V.bool_value)
result.add_check("guidance_scale", batch.guidance_scale,
V.positive_float)
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify conditioning stage outputs."""
result = VerificationResult()
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
return result
+19
View File
@@ -11,6 +11,8 @@ from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
from fastvideo.v1.pipelines.stages.validators import VerificationResult
from fastvideo.v1.utils import PRECISION_TO_TYPE
logger = init_logger(__name__)
@@ -27,6 +29,23 @@ class DecodingStage(PipelineStage):
def __init__(self, vae) -> None:
self.vae: ParallelTiledVAE = vae
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify decoding stage inputs."""
result = VerificationResult()
# Denoised latents for VAE decoding: [batch_size, channels, frames, height_latents, width_latents]
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify decoding stage outputs."""
result = VerificationResult()
# Decoded video/images: [batch_size, channels, frames, height, width]
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
return result
def forward(
self,
batch: ForwardBatch,
+59 -33
View File
@@ -3,15 +3,15 @@
Denoising stage for diffusion pipelines.
"""
import importlib.util
import inspect
from typing import Any, Dict, Iterable, List, Optional
from typing import Any, Dict, Iterable, Optional
import torch
from einops import rearrange
from tqdm.auto import tqdm
from fastvideo.v1.attention import get_attn_backend
from fastvideo.v1.configs.pipelines.base import STA_Mode
from fastvideo.v1.distributed import (get_sp_parallel_rank, get_sp_world_size,
get_torch_device, get_world_group)
from fastvideo.v1.distributed.communication_op import (
@@ -21,19 +21,24 @@ from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
from fastvideo.v1.pipelines.stages.validators import VerificationResult
from fastvideo.v1.platforms import AttentionBackendEnum
from fastvideo.v1.utils import dict_to_3d_list
st_attn_available = False
if importlib.util.find_spec("st_attn") is not None:
st_attn_available = True
try:
from fastvideo.v1.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
st_attn_available = False
vsa_available = False
if importlib.util.find_spec("vsa") is not None:
vsa_available = True
try:
from fastvideo.v1.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
vsa_available = False
logger = init_logger(__name__)
@@ -117,20 +122,6 @@ class DenoisingStage(PipelineStage):
num_warmup_steps = len(
timesteps) - num_inference_steps * self.scheduler.order
# Create 3D list for mask strategy
def dict_to_3d_list(mask_strategy,
t_max=50,
l_max=60,
h_max=24) -> List:
result = [[[None for _ in range(h_max)] for _ in range(l_max)]
for _ in range(t_max)]
if mask_strategy is None:
return result
for key, value in mask_strategy.items():
t, layer, h = map(int, key.split('_'))
result[t][layer][h] = value
return result
# Prepare image latents and embeddings for I2V generation
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
@@ -143,7 +134,8 @@ class DenoisingStage(PipelineStage):
self.transformer.forward,
{
"encoder_hidden_states_image": image_embeds,
"mask_strategy": dict_to_3d_list(None)
"mask_strategy": dict_to_3d_list(
None, t_max=50, l_max=60, h_max=24)
},
)
@@ -304,7 +296,7 @@ class DenoisingStage(PipelineStage):
batch.latents = latents
# Save STA mask search results if needed
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == 'STA_searching':
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING:
self.save_sta_search_results(batch)
if fastvideo_args.use_cpu_offload:
@@ -402,7 +394,7 @@ class DenoisingStage(PipelineStage):
raise NotImplementedError(
"STA mask search/tuning is not supported for this resolution")
if STA_mode == "STA_searching" or STA_mode == "STA_tuning" or STA_mode == "STA_tuning_cfg":
if STA_mode == STA_Mode.STA_SEARCHING or STA_mode == STA_Mode.STA_TUNING or STA_mode == STA_Mode.STA_TUNING_CFG:
size = (batch.width, batch.height)
if size == (1280, 768):
# TODO: make it configurable
@@ -424,18 +416,18 @@ class DenoisingStage(PipelineStage):
layer_num += self.transformer.config.num_single_layers
head_num = self.transformer.config.num_attention_heads
if STA_mode == "STA_searching":
if STA_mode == STA_Mode.STA_SEARCHING:
STA_param = configure_sta(
mode='STA_searching',
mode=STA_Mode.STA_SEARCHING,
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
mask_candidates=sparse_mask_candidates_searching +
full_mask, # last is full mask; Can add more sparse masks while keep last one as full mask
)
elif STA_mode == 'STA_tuning':
elif STA_mode == STA_Mode.STA_TUNING:
STA_param = configure_sta(
mode='STA_tuning',
mode=STA_Mode.STA_TUNING,
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
@@ -448,9 +440,9 @@ class DenoisingStage(PipelineStage):
save_dir=
f'output/mask_search_strategy_{size[0]}x{size[1]}/', # Custom save directory
timesteps=timesteps_num)
elif STA_mode == 'STA_tuning_cfg':
elif STA_mode == STA_Mode.STA_TUNING_CFG:
STA_param = configure_sta(
mode='STA_tuning_cfg',
mode=STA_Mode.STA_TUNING_CFG,
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
@@ -463,12 +455,12 @@ class DenoisingStage(PipelineStage):
skip_time_steps=skip_time_steps,
save_dir=f'output/mask_search_strategy_{size[0]}x{size[1]}/',
timesteps=timesteps_num)
elif STA_mode == 'STA_inference':
elif STA_mode == STA_Mode.STA_INFERENCE:
import fastvideo.v1.envs as envs
config_file = envs.FASTVIDEO_ATTENTION_CONFIG
if config_file is None:
raise ValueError("FASTVIDEO_ATTENTION_CONFIG is not set")
STA_param = configure_sta(mode='STA_inference',
STA_param = configure_sta(mode=STA_Mode.STA_INFERENCE,
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
@@ -517,3 +509,37 @@ class DenoisingStage(PipelineStage):
mask_strategies=sparse_mask_candidates_searching,
output_dir=f'output/mask_search_result_neg_{size[0]}x{size[1]}/'
)
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify denoising stage inputs."""
result = VerificationResult()
result.add_check("timesteps", batch.timesteps,
[V.is_tensor, V.min_dims(1)])
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
result.add_check("image_embeds", batch.image_embeds, V.is_list)
result.add_check("image_latent", batch.image_latent,
V.none_or_tensor_with_dims(5))
result.add_check("num_inference_steps", batch.num_inference_steps,
V.positive_int)
result.add_check("guidance_scale", batch.guidance_scale,
V.positive_float)
result.add_check("eta", batch.eta, V.non_negative_float)
result.add_check("generator", batch.generator,
V.generator_or_list_generators)
result.add_check("do_classifier_free_guidance",
batch.do_classifier_free_guidance, V.bool_value)
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify denoising stage outputs."""
result = VerificationResult()
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
return result
+40 -14
View File
@@ -12,10 +12,12 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.models.vision_utils import (get_default_height_width,
load_image, normalize,
numpy_to_pt, pil_to_numpy, resize)
normalize, numpy_to_pt,
pil_to_numpy, resize)
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import V # Import validators
from fastvideo.v1.pipelines.stages.validators import VerificationResult
from fastvideo.v1.utils import PRECISION_TO_TYPE
logger = init_logger(__name__)
@@ -49,22 +51,26 @@ class EncodingStage(PipelineStage):
"""
self.vae = self.vae.to(get_torch_device())
image_path = batch.image_path
# TODO(will): remove this once we add input/output validation for stages
if image_path is None:
raise ValueError("Image Path must be provided")
assert batch.height is not None
assert batch.width is not None
latent_height = batch.height // self.vae.spatial_compression_ratio
latent_width = batch.width // self.vae.spatial_compression_ratio
image = load_image(image_path)
image = self.preprocess(
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=batch.height,
width=batch.width).to(get_torch_device(), dtype=torch.float32)
image = image.unsqueeze(2)
image = batch.preprocessed_image
# TODO(will)
if image is None:
assert batch.pil_image is not None
image = batch.pil_image
image = self.preprocess(
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=batch.height,
width=batch.width).to(get_torch_device(), dtype=torch.float32)
image = image.unsqueeze(2)
else:
image = image.transpose(1, 2)
logger.info("image: %s", image.shape)
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1],
@@ -95,7 +101,7 @@ class EncodingStage(PipelineStage):
generator = batch.generator
if generator is None:
raise ValueError("Generator must be provided")
latent_condition = self.retrieve_latents(encoder_output, generator[0])
latent_condition = self.retrieve_latents(encoder_output, generator)
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")
@@ -174,3 +180,23 @@ class EncodingStage(PipelineStage):
image = normalize(image)
return image
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify encoding stage inputs."""
result = VerificationResult()
# result.add_check("pil_image", batch.pil_image)
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
result.add_check("generator", batch.generator,
V.generator_or_list_generators)
result.add_check("num_frames", batch.num_frames, V.positive_int)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify encoding stage outputs."""
result = VerificationResult()
result.add_check("image_latent", batch.image_latent,
[V.is_tensor, V.with_dims(5)])
return result
@@ -11,9 +11,10 @@ from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vision_utils import load_image
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
from fastvideo.v1.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
@@ -56,7 +57,7 @@ class ImageEncodingStage(PipelineStage):
if fastvideo_args.use_cpu_offload:
self.image_encoder = self.image_encoder.to(get_torch_device())
image = load_image(batch.image_path)
image = batch.pil_image
image_inputs = self.image_processor(
images=image, return_tensors="pt").to(get_torch_device())
@@ -71,3 +72,19 @@ class ImageEncodingStage(PipelineStage):
torch.cuda.empty_cache()
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify image encoding stage inputs."""
result = VerificationResult()
result.add_check("pil_image", batch.pil_image, V.not_none)
result.add_check("image_embeds", batch.image_embeds, V.is_list)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify image encoding stage outputs."""
result = VerificationResult()
result.add_check("image_embeds", batch.image_embeds,
V.list_of_tensors_dims(3))
return result
@@ -7,11 +7,17 @@ import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vision_utils import load_image
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import (StageValidators,
VerificationResult)
logger = init_logger(__name__)
# Alias for convenience
V = StageValidators
class InputValidationStage(PipelineStage):
"""
@@ -86,4 +92,37 @@ class InputValidationStage(PipelineStage):
f"Guidance scale must be positive, but got {batch.guidance_scale}"
)
# for i2v, get image from image_path
if batch.image_path is not None:
image = load_image(batch.image_path)
batch.pil_image = image
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify input validation stage inputs."""
result = VerificationResult()
result.add_check("seed", batch.seed, [V.not_none, V.positive_int])
result.add_check("num_videos_per_prompt", batch.num_videos_per_prompt,
V.positive_int)
result.add_check(
"prompt_or_embeds", None, lambda _: V.string_or_list_strings(
batch.prompt) or V.list_not_empty(batch.prompt_embeds))
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
result.add_check("num_inference_steps", batch.num_inference_steps,
V.positive_int)
result.add_check(
"guidance_scale", batch.guidance_scale, lambda x: not batch.
do_classifier_free_guidance or V.positive_float(x))
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify input validation stage outputs."""
result = VerificationResult()
result.add_check("seeds", batch.seeds, V.list_not_empty)
result.add_check("generator", batch.generator,
V.generator_or_list_generators)
return result
@@ -2,6 +2,7 @@
"""
Latent preparation stage for diffusion pipelines.
"""
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.v1.distributed import get_torch_device
@@ -9,6 +10,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
from fastvideo.v1.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
@@ -126,4 +129,32 @@ class LatentPreparationStage(PipelineStage):
latent_num_frames = (video_length - 1) // temporal_scale_factor + 1
else: # stepvideo only
latent_num_frames = video_length // 17 * 3
return latent_num_frames
return int(latent_num_frames)
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify latent preparation stage inputs."""
result = VerificationResult()
result.add_check(
"prompt_or_embeds", None, lambda _: V.string_or_list_strings(
batch.prompt) or V.list_not_empty(batch.prompt_embeds))
result.add_check("prompt_embeds", batch.prompt_embeds,
V.list_of_tensors)
result.add_check("num_videos_per_prompt", batch.num_videos_per_prompt,
V.positive_int)
result.add_check("generator", batch.generator,
V.generator_or_list_generators)
result.add_check("num_frames", batch.num_frames, V.positive_int)
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
result.add_check("latents", batch.latents, V.none_or_tensor)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify latent preparation stage outputs."""
result = VerificationResult()
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
result.add_check("raw_latent_shape", batch.raw_latent_shape, V.is_tuple)
return result
@@ -1,10 +1,14 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
from fastvideo.v1.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
@@ -47,3 +51,29 @@ class StepvideoPromptEncodingStage(PipelineStage):
batch.clip_embedding_pos = pos_clip
batch.clip_embedding_neg = neg_clip
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify stepvideo encoding stage inputs."""
result = VerificationResult()
result.add_check("prompt", batch.prompt, V.string_not_empty)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify stepvideo encoding stage outputs."""
result = VerificationResult()
result.add_check("prompt_embeds", batch.prompt_embeds,
[V.is_tensor, V.with_dims(3)])
result.add_check("negative_prompt_embeds", batch.negative_prompt_embeds,
[V.is_tensor, V.with_dims(3)])
result.add_check("prompt_attention_mask", batch.prompt_attention_mask,
[V.is_tensor, V.with_dims(2)])
result.add_check("negative_attention_mask",
batch.negative_attention_mask,
[V.is_tensor, V.with_dims(2)])
result.add_check("clip_embedding_pos", batch.clip_embedding_pos,
[V.is_tensor, V.with_dims(2)])
result.add_check("clip_embedding_neg", batch.clip_embedding_neg,
[V.is_tensor, V.with_dims(2)])
return result
@@ -12,6 +12,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
from fastvideo.v1.pipelines.stages.validators import VerificationResult
logger = (__name__)
@@ -113,3 +115,30 @@ class TextEncodingStage(PipelineStage):
torch.cuda.empty_cache()
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify text encoding stage inputs."""
result = VerificationResult()
result.add_check("prompt", batch.prompt, V.string_or_list_strings)
result.add_check(
"negative_prompt", batch.negative_prompt, lambda x: not batch.
do_classifier_free_guidance or V.string_not_empty(x))
result.add_check("do_classifier_free_guidance",
batch.do_classifier_free_guidance, V.bool_value)
result.add_check("prompt_embeds", batch.prompt_embeds, V.is_list)
result.add_check("negative_prompt_embeds", batch.negative_prompt_embeds,
V.none_or_list)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify text encoding stage outputs."""
result = VerificationResult()
result.add_check("prompt_embeds", batch.prompt_embeds,
V.list_of_tensors_min_dims(2))
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds,
lambda x: not batch.do_classifier_free_guidance or V.
list_of_tensors_with_min_dims(x, 2))
return result
@@ -12,6 +12,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
from fastvideo.v1.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
@@ -95,3 +97,22 @@ class TimestepPreparationStage(PipelineStage):
batch.timesteps = timesteps
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify timestep preparation stage inputs."""
result = VerificationResult()
result.add_check("num_inference_steps", batch.num_inference_steps,
V.positive_int)
result.add_check("timesteps", batch.timesteps, V.none_or_tensor)
result.add_check("sigmas", batch.sigmas, V.none_or_list)
result.add_check("n_tokens", batch.n_tokens, V.none_or_positive_int)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify timestep preparation stage outputs."""
result = VerificationResult()
result.add_check("timesteps", batch.timesteps,
[V.is_tensor, V.with_dims(1)])
return result
+486
View File
@@ -0,0 +1,486 @@
# SPDX-License-Identifier: Apache-2.0
"""
Common validators for pipeline stage verification.
This module provides reusable validation functions that can be used across
all pipeline stages for input/output verification.
"""
from typing import Any, Callable, Dict, List, Optional, Union
import torch
class StageValidators:
"""Common validators for pipeline stages."""
@staticmethod
def not_none(value: Any) -> bool:
"""Check if value is not None."""
return value is not None
@staticmethod
def positive_int(value: Any) -> bool:
"""Check if value is a positive integer."""
return isinstance(value, int) and value > 0
@staticmethod
def positive_float(value: Any) -> bool:
"""Check if value is a positive float."""
return isinstance(value, (int, float)) and value > 0
@staticmethod
def non_negative_float(value: Any) -> bool:
"""Check if value is a non-negative float."""
return isinstance(value, (int, float)) and value >= 0
@staticmethod
def divisible_by(value: Any, divisor: int) -> bool:
"""Check if value is divisible by divisor."""
return value is not None and isinstance(value,
int) and value % divisor == 0
@staticmethod
def is_tensor(value: Any) -> bool:
"""Check if value is a torch tensor and doesn't contain NaN values."""
if not isinstance(value, torch.Tensor):
return False
return not torch.isnan(value).any().item()
@staticmethod
def tensor_with_dims(value: Any, dims: int) -> bool:
"""Check if value is a tensor with specific dimensions and no NaN values."""
if not isinstance(value, torch.Tensor):
return False
if value.dim() != dims:
return False
return not torch.isnan(value).any().item()
@staticmethod
def tensor_min_dims(value: Any, min_dims: int) -> bool:
"""Check if value is a tensor with at least min_dims dimensions and no NaN values."""
if not isinstance(value, torch.Tensor):
return False
if value.dim() < min_dims:
return False
return not torch.isnan(value).any().item()
@staticmethod
def tensor_shape_matches(value: Any, expected_shape: tuple) -> bool:
"""Check if tensor shape matches expected shape (None for any size) and no NaN values."""
if not isinstance(value, torch.Tensor):
return False
if len(value.shape) != len(expected_shape):
return False
for actual, expected in zip(value.shape, expected_shape):
if expected is not None and actual != expected:
return False
return not torch.isnan(value).any().item()
@staticmethod
def list_not_empty(value: Any) -> bool:
"""Check if value is a non-empty list."""
return isinstance(value, list) and len(value) > 0
@staticmethod
def list_length(value: Any, length: int) -> bool:
"""Check if list has specific length."""
return isinstance(value, list) and len(value) == length
@staticmethod
def list_min_length(value: Any, min_length: int) -> bool:
"""Check if list has at least min_length items."""
return isinstance(value, list) and len(value) >= min_length
@staticmethod
def string_not_empty(value: Any) -> bool:
"""Check if value is a non-empty string."""
return isinstance(value, str) and len(value.strip()) > 0
@staticmethod
def string_or_list_strings(value: Any) -> bool:
"""Check if value is a string or list of strings."""
if isinstance(value, str):
return True
if isinstance(value, list):
return all(isinstance(item, str) for item in value)
return False
@staticmethod
def bool_value(value: Any) -> bool:
"""Check if value is a boolean."""
return isinstance(value, bool)
@staticmethod
def generator_or_list_generators(value: Any) -> bool:
"""Check if value is a Generator or list of Generators."""
if isinstance(value, torch.Generator):
return True
if isinstance(value, list):
return all(isinstance(item, torch.Generator) for item in value)
return False
@staticmethod
def is_list(value: Any) -> bool:
"""Check if value is a list (can be empty)."""
return isinstance(value, list)
@staticmethod
def is_tuple(value: Any) -> bool:
"""Check if value is a tuple."""
return isinstance(value, tuple)
@staticmethod
def none_or_tensor(value: Any) -> bool:
"""Check if value is None or a tensor without NaN values."""
if value is None:
return True
if not isinstance(value, torch.Tensor):
return False
return not torch.isnan(value).any().item()
@staticmethod
def list_of_tensors_with_dims(value: Any, dims: int) -> bool:
"""Check if value is a non-empty list where all items are tensors with specific dimensions and no NaN values."""
if not isinstance(value, list) or len(value) == 0:
return False
for item in value:
if not isinstance(item, torch.Tensor):
return False
if item.dim() != dims:
return False
if torch.isnan(item).any().item():
return False
return True
@staticmethod
def list_of_tensors(value: Any) -> bool:
"""Check if value is a non-empty list where all items are tensors without NaN values."""
if not isinstance(value, list) or len(value) == 0:
return False
for item in value:
if not isinstance(item, torch.Tensor):
return False
if torch.isnan(item).any().item():
return False
return True
@staticmethod
def list_of_tensors_with_min_dims(value: Any, min_dims: int) -> bool:
"""Check if value is a non-empty list where all items are tensors with at least min_dims dimensions and no NaN values."""
if not isinstance(value, list) or len(value) == 0:
return False
for item in value:
if not isinstance(item, torch.Tensor):
return False
if item.dim() < min_dims:
return False
if torch.isnan(item).any().item():
return False
return True
@staticmethod
def none_or_tensor_with_dims(dims: int) -> Callable[[Any], bool]:
"""Return a validator that checks if value is None or a tensor with specific dimensions and no NaN values."""
def validator(value: Any) -> bool:
if value is None:
return True
if not isinstance(value, torch.Tensor):
return False
if value.dim() != dims:
return False
return not torch.isnan(value).any().item()
return validator
@staticmethod
def none_or_list(value: Any) -> bool:
"""Check if value is None or a list."""
return value is None or isinstance(value, list)
@staticmethod
def none_or_positive_int(value: Any) -> bool:
"""Check if value is None or a positive integer."""
return value is None or (isinstance(value, int) and value > 0)
# Helper methods that return functions for common patterns
@staticmethod
def with_dims(dims: int) -> Callable[[Any], bool]:
"""Return a validator that checks if tensor has specific dimensions and no NaN values."""
def validator(value: Any) -> bool:
return StageValidators.tensor_with_dims(value, dims)
return validator
@staticmethod
def min_dims(min_dims: int) -> Callable[[Any], bool]:
"""Return a validator that checks if tensor has at least min_dims dimensions and no NaN values."""
def validator(value: Any) -> bool:
return StageValidators.tensor_min_dims(value, min_dims)
return validator
@staticmethod
def divisible(divisor: int) -> Callable[[Any], bool]:
"""Return a validator that checks if value is divisible by divisor."""
def validator(value: Any) -> bool:
return StageValidators.divisible_by(value, divisor)
return validator
@staticmethod
def positive_int_divisible(divisor: int) -> Callable[[Any], bool]:
"""Return a validator that checks if value is a positive integer divisible by divisor."""
def validator(value: Any) -> bool:
return (isinstance(value, int) and value > 0
and StageValidators.divisible_by(value, divisor))
return validator
@staticmethod
def list_of_tensors_dims(dims: int) -> Callable[[Any], bool]:
"""Return a validator that checks if value is a list of tensors with specific dimensions and no NaN values."""
def validator(value: Any) -> bool:
return StageValidators.list_of_tensors_with_dims(value, dims)
return validator
@staticmethod
def list_of_tensors_min_dims(min_dims: int) -> Callable[[Any], bool]:
"""Return a validator that checks if value is a list of tensors with at least min_dims dimensions and no NaN values."""
def validator(value: Any) -> bool:
return StageValidators.list_of_tensors_with_min_dims(
value, min_dims)
return validator
class ValidationFailure:
"""Details about a specific validation failure."""
def __init__(self,
validator_name: str,
actual_value: Any,
expected: Optional[str] = None,
error_msg: Optional[str] = None):
self.validator_name = validator_name
self.actual_value = actual_value
self.expected = expected
self.error_msg = error_msg
def __str__(self) -> str:
parts = [f"Validator '{self.validator_name}' failed"]
if self.error_msg:
parts.append(f"Error: {self.error_msg}")
# Add actual value info (but limit very long representations)
actual_str = self._format_value(self.actual_value)
parts.append(f"Actual: {actual_str}")
if self.expected:
parts.append(f"Expected: {self.expected}")
return ". ".join(parts)
def _format_value(self, value: Any) -> str:
"""Format a value for display in error messages."""
if value is None:
return "None"
elif isinstance(value, torch.Tensor):
return f"tensor(shape={list(value.shape)}, dtype={value.dtype})"
elif isinstance(value, list):
if len(value) == 0:
return "[]"
elif len(value) <= 3:
item_strs = [self._format_value(item) for item in value]
return f"[{', '.join(item_strs)}]"
else:
return f"list(length={len(value)}, first_item={self._format_value(value[0])})"
elif isinstance(value, str):
if len(value) > 50:
return f"'{value[:47]}...'"
else:
return f"'{value}'"
else:
return f"{type(value).__name__}({value})"
class VerificationResult:
"""Wrapper class for stage verification results."""
def __init__(self) -> None:
self._checks: Dict[str, bool] = {}
self._failures: Dict[str, List[ValidationFailure]] = {}
def add_check(
self, field_name: str, value: Any,
validators: Union[Callable[[Any], bool], List[Callable[[Any], bool]]]
) -> 'VerificationResult':
"""
Add a validation check for a field.
Args:
field_name: Name of the field being checked
value: The actual value to validate
validators: Single validation function or list of validation functions.
Each function will be called with the value as its first argument.
Returns:
Self for method chaining
Examples:
# Single validator
result.add_check("tensor", my_tensor, V.is_tensor)
# Multiple validators (all must pass)
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
# Using partial functions for parameters
result.add_check("height", batch.height, [V.not_none, V.divisible(8)])
"""
if not isinstance(validators, list):
validators = [validators]
failures = []
all_passed = True
# Apply all validators and collect detailed failure info
for validator in validators:
try:
passed = validator(value)
if not passed:
all_passed = False
failure = self._create_validation_failure(validator, value)
failures.append(failure)
except Exception as e:
# If any validator raises an exception, consider the check failed
all_passed = False
validator_name = getattr(validator, '__name__', str(validator))
failure = ValidationFailure(
validator_name=validator_name,
actual_value=value,
error_msg=f"Exception during validation: {str(e)}")
failures.append(failure)
self._checks[field_name] = all_passed
if not all_passed:
self._failures[field_name] = failures
return self
def _create_validation_failure(self, validator: Callable,
value: Any) -> ValidationFailure:
"""Create a ValidationFailure with detailed information."""
validator_name = getattr(validator, '__name__', str(validator))
# Try to extract meaningful expected value info based on validator type
expected = None
error_msg = None
# Handle common validator patterns
if hasattr(validator, '__closure__') and validator.__closure__:
# This is likely a closure (like our helper functions)
if 'dims' in validator_name or 'with_dims' in str(validator):
if isinstance(value, torch.Tensor):
expected = f"tensor with {validator.__closure__[0].cell_contents} dimensions"
else:
expected = "tensor with specific dimensions"
elif 'divisible' in str(validator):
expected = f"integer divisible by {validator.__closure__[0].cell_contents}"
# Handle specific validator types and check for NaN values
if validator_name == 'is_tensor':
expected = "torch.Tensor without NaN values"
if isinstance(value,
torch.Tensor) and torch.isnan(value).any().item():
error_msg = f"tensor contains {torch.isnan(value).sum().item()} NaN values"
elif validator_name == 'positive_int':
expected = "positive integer"
elif validator_name == 'not_none':
expected = "non-None value"
elif validator_name == 'list_not_empty':
expected = "non-empty list"
elif validator_name == 'bool_value':
expected = "boolean value"
elif 'tensor_with_dims' in validator_name or 'tensor_min_dims' in validator_name:
if isinstance(value, torch.Tensor):
if torch.isnan(value).any().item():
error_msg = f"tensor has {value.dim()} dimensions but contains {torch.isnan(value).sum().item()} NaN values"
else:
error_msg = f"tensor has {value.dim()} dimensions"
elif validator_name == 'is_list':
expected = "list"
elif validator_name == 'none_or_tensor':
expected = "None or tensor without NaN values"
if isinstance(value,
torch.Tensor) and torch.isnan(value).any().item():
error_msg = f"tensor contains {torch.isnan(value).sum().item()} NaN values"
elif validator_name == 'list_of_tensors':
expected = "non-empty list of tensors without NaN values"
if isinstance(value, list) and len(value) > 0:
nan_count = 0
for item in value:
if isinstance(
item,
torch.Tensor) and torch.isnan(item).any().item():
nan_count += torch.isnan(item).sum().item()
if nan_count > 0:
error_msg = f"list contains tensors with total {nan_count} NaN values"
elif 'list_of_tensors_with_dims' in validator_name:
expected = "non-empty list of tensors with specific dimensions and no NaN values"
if isinstance(value, list) and len(value) > 0:
nan_count = 0
for item in value:
if isinstance(
item,
torch.Tensor) and torch.isnan(item).any().item():
nan_count += torch.isnan(item).sum().item()
if nan_count > 0:
error_msg = f"list contains tensors with total {nan_count} NaN values"
return ValidationFailure(validator_name=validator_name,
actual_value=value,
expected=expected,
error_msg=error_msg)
def is_valid(self) -> bool:
"""Check if all validations passed."""
return all(self._checks.values())
def get_failed_fields(self) -> List[str]:
"""Get list of fields that failed validation."""
return [field for field, passed in self._checks.items() if not passed]
def get_detailed_failures(self) -> Dict[str, List[ValidationFailure]]:
"""Get detailed failure information for each failed field."""
return self._failures.copy()
def get_failure_summary(self) -> str:
"""Get a comprehensive summary of all validation failures."""
if self.is_valid():
return "All validations passed"
summary_parts = []
for field_name, failures in self._failures.items():
field_summary = f"\n Field '{field_name}':"
for i, failure in enumerate(failures, 1):
field_summary += f"\n {i}. {failure}"
summary_parts.append(field_summary)
return "Validation failures:" + "".join(summary_parts)
def to_dict(self) -> dict:
"""Convert to dictionary for backward compatibility."""
return self._checks.copy()
# Alias for convenience
V = StageValidators
@@ -76,4 +76,36 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
stage=DecodingStage(vae=self.get_module("vae")))
class WanImageToVideoValidationPipeline(ComposedPipelineBase):
"""
I2V Validation pipeline for Wan2.1, assumes that the input are preprocess latents.
"""
_required_config_modules = ["vae", "scheduler", "transformer"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
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")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanImageToVideoPipeline
+10 -18
View File
@@ -123,14 +123,13 @@ class CudaPlatformBase(Platform):
SlidingTileAttentionBackend)
logger.info("Using Sliding Tile Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "SLIDING_TILE_ATTN"
return "fastvideo.v1.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info(
"Sliding Tile Attention backend is not installed. Fall back to Flash Attention."
)
logger.error(
"Failed to import Sliding Tile Attention backend: %s",
str(e))
raise ImportError(
"Sliding Tile Attention backend is not installed. ") from e
elif selected_backend == AttentionBackendEnum.SAGE_ATTN:
try:
from sageattention import sageattn # noqa: F401
@@ -139,8 +138,6 @@ class CudaPlatformBase(Platform):
SageAttentionBackend)
logger.info("Using Sage Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "SAGE_ATTN"
return "fastvideo.v1.attention.backends.sage_attn.SageAttentionBackend"
except ImportError as e:
logger.info(e)
@@ -155,14 +152,13 @@ class CudaPlatformBase(Platform):
VideoSparseAttentionBackend)
logger.info("Using Video Sparse Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "VIDEO_SPARSE_ATTN"
return "fastvideo.v1.attention.backends.video_sparse_attn.VideoSparseAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info(
"Video Sparse Attention backend is not installed. Fall back to Flash Attention."
)
logger.error(
"Failed to import Video Sparse Attention backend: %s",
str(e))
raise ImportError(
"Video Sparse Attention backend is not installed. ") from e
elif selected_backend == AttentionBackendEnum.TORCH_SDPA:
logger.info("Using Torch SDPA backend.")
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
@@ -209,14 +205,10 @@ class CudaPlatformBase(Platform):
if target_backend == AttentionBackendEnum.TORCH_SDPA:
logger.info("Using Torch SDPA backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "TORCH_SDPA"
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
logger.info("Using Flash Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "FLASH_ATTN"
return "fastvideo.v1.attention.backends.flash_attn.FlashAttentionBackend"
@classmethod
+1
View File
@@ -169,6 +169,7 @@ class Platform:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
@classmethod
def verify_model_arch(cls, model_arch: str) -> None:
@@ -0,0 +1,52 @@
import os
import sys
import subprocess
from pathlib import Path
import pytest
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "1"
# Set environment variables
os.environ["FASTVIDEO_ATTENTION_CONFIG"] = "assets/mask_strategy_wan.json"
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLIDING_TILE_ATTN"
def test_inference():
"""Test the inference functionality"""
# Create command as in wan_14B-STA.sh
cmd = [
"fastvideo", "generate",
"--model-path", "Wan-AI/Wan2.1-T2V-14B-Diffusers",
"--sp-size", "1",
"--tp-size", "1",
"--num-gpus", "1",
"--height", "768",
"--width", "1280",
"--num-frames", "69",
"--num-inference-steps", "2",
"--fps", "16",
"--guidance-scale", "5.0",
"--flow-shift", "5.0",
"--prompt", "A majestic lion strides across the golden savanna, its powerful frame glistening under the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's commanding presence. The tone is vibrant, embodying the raw energy of the wild. Low angle, steady tracking shot, cinematic.",
"--negative-prompt", "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
"--seed", "1024",
"--output-path", "outputs_video/STA_1024/",
]
# Run the command
subprocess.run(cmd, check=True)
# Verify output directory exists
output_dir = Path("outputs_video/STA_1024/")
assert output_dir.exists(), f"Output directory {output_dir} does not exist"
# Verify that video files were generated
video_files = list(output_dir.glob("*.mp4"))
assert len(video_files) > 0, "No video files were generated"
# Verify the video file properties
for video_file in video_files:
assert video_file.stat().st_size > 0, f"Video file {video_file} is empty"
if __name__ == "__main__":
test_inference()
+98
View File
@@ -0,0 +1,98 @@
import modal
app = modal.App()
import os
image_version = os.getenv("IMAGE_VERSION", "latest")
image_tag = f"ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:{image_version}"
print(f"Using image: {image_tag}")
image = (
modal.Image.from_registry(image_tag, add_python="3.12")
.apt_install("cmake", "pkg-config", "build-essential", "curl", "libssl-dev")
.run_commands("curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain stable")
.run_commands("echo 'source ~/.cargo/env' >> ~/.bashrc")
.env({"PATH": "/root/.cargo/bin:$PATH"})
.run_commands("/bin/bash -c 'source $HOME/.local/bin/env && source /opt/venv/bin/activate && cd /FastVideo && uv pip install -e .[test]'")
)
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_encoder_tests():
"""Run encoder tests on L40S GPU"""
import subprocess
import sys
import os
os.chdir("/FastVideo")
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/encoders -s
"""
result = subprocess.run([
"/bin/bash", "-c", command
], stdout=sys.stdout, stderr=sys.stderr, check=False)
sys.exit(result.returncode)
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_vae_tests():
"""Run VAE tests on L40S GPU"""
import subprocess
import sys
import os
os.chdir("/FastVideo")
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/vaes -s
"""
result = subprocess.run([
"/bin/bash", "-c", command
], stdout=sys.stdout, stderr=sys.stderr, check=False)
sys.exit(result.returncode)
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_transformer_tests():
"""Run transformer tests on L40S GPU"""
import subprocess
import sys
import os
os.chdir("/FastVideo")
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/transformers -s
"""
result = subprocess.run([
"/bin/bash", "-c", command
], stdout=sys.stdout, stderr=sys.stderr, check=False)
sys.exit(result.returncode)
@app.function(gpu="L40S:2", image=image, timeout=3600)
def run_ssim_tests():
"""Run SSIM tests on 2x L40S GPUs"""
import subprocess
import sys
import os
os.chdir("/FastVideo")
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/ssim -vs
"""
result = subprocess.run([
"/bin/bash", "-c", command
], stdout=sys.stdout, stderr=sys.stderr, check=False)
sys.exit(result.returncode)
@@ -0,0 +1 @@
{"step_time":2.245914653001819,"_wandb":{"runtime":1434},"learning_rate":1e-05,"grad_norm":0.57421875,"avg_step_time":1.1814782944297622,"train_loss":0.07932619750499725,"vsa_sparsity":0,"_timestamp":1.750578625921253e+09,"validation_videos_40_steps":{"count":1,"videos":[{"size":420969,"path":"media/videos/validation_videos_40_steps_900_581ff5eae2909d3a7b36.mp4","_type":"video-file","sha256":"581ff5eae2909d3a7b362dcb24d060c006c09e4d4deb44b82f4aa697f6789ba7"}],"captions":false,"_type":"videos"},"_runtime":1434.62395329,"_step":901}
@@ -0,0 +1,177 @@
import os
from pathlib import Path
from huggingface_hub import snapshot_download
import shutil
import subprocess
import sys
from fastvideo.v1.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
# Import the training pipeline
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
NUM_NODES = "1"
MODEL_PATH = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
# preprocessing
DATA_DIR = "data"
LOCAL_RAW_DATA_DIR = Path(os.path.join(DATA_DIR, "cats"))
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
PREPROCESSING_ENTRY_FILE_PATH = "fastvideo/v1/pipelines/preprocess/v1_preprocess.py"
LOCAL_PREPROCESSED_DATA_DIR = Path(os.path.join(DATA_DIR, "cats_preprocessed_data_i2v"))
# training
NUM_GPUS_PER_NODE_TRAINING = "8"
TRAINING_ENTRY_FILE_PATH = "fastvideo/v1/training/wan_i2v_training_pipeline.py"
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "combined_parquet_dataset")
LOCAL_VALIDATION_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "validation_parquet_dataset")
LOCAL_OUTPUT_DIR = Path(os.path.join(DATA_DIR, "outputs"))
def download_data():
# create the data dir if it doesn't exist
data_dir = Path(DATA_DIR)
# if data_dir.exists():
# print(f"Removing existing data directory at {data_dir}")
# shutil.rmtree(data_dir)
print(f"Creating data directory at {data_dir}")
os.makedirs(data_dir)
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
try:
# result = snapshot_download(
# repo_id="wlsaidhi/cats-overfit-merged",
# local_dir=str(LOCAL_RAW_DATA_DIR),
# repo_type="dataset",
# resume_download=True,
# token=os.environ.get("HF_TOKEN"), # In case authentication is needed
# )
print(f"Download completed successfully. Files downloaded to: {result}")
# Verify the download
if not LOCAL_RAW_DATA_DIR.exists():
raise RuntimeError(f"Download appeared to succeed but {LOCAL_RAW_DATA_DIR} does not exist")
# List downloaded files
print("Downloaded files:")
for file in LOCAL_RAW_DATA_DIR.rglob("*"):
if file.is_file():
print(f" - {file.relative_to(LOCAL_RAW_DATA_DIR)}")
except Exception as e:
print(f"Error during download: {str(e)}")
raise
def run_preprocessing():
# Run torchrun command
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
PREPROCESSING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--data_merge_path", os.path.join(LOCAL_RAW_DATA_DIR, "merge_1_sample.txt"),
"--preprocess_video_batch_size", "1",
"--max_height", "480",
"--max_width", "832",
"--num_frames", "77",
"--dataloader_num_workers", "0",
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
"--train_fps", "16",
"--validation_dataset_file", os.path.join(LOCAL_RAW_DATA_DIR, "validation_i2v_prompt_1_sample.json"),
"--samples_per_file", "1",
"--flush_frequency", "1",
"--video_length_tolerance_range", "5",
"--preprocess_task", "i2v",
]
process = subprocess.run(cmd, check=True)
def run_training():
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_TRAINING,
TRAINING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--inference_mode", "False",
"--pretrained_model_name_or_path", MODEL_PATH,
"--data_path", LOCAL_TRAINING_DATA_DIR,
"--validation_preprocessed_path", LOCAL_VALIDATION_DATA_DIR,
"--train_batch_size", "1",
"--num_latent_t", "8",
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
"--sp_size", NUM_GPUS_PER_NODE_TRAINING,
"--tp_size", NUM_GPUS_PER_NODE_TRAINING,
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", NUM_GPUS_PER_NODE_TRAINING,
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "10",
"--gradient_accumulation_steps", "1",
"--max_train_steps", "901",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--checkpointing_steps", "6000",
"--validation_steps", "100",
"--validation_sampling_steps", "40",
"--log_validation",
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--cfg", "0.0",
"--output_dir", LOCAL_OUTPUT_DIR,
"--tracker_project_name", "wan_i2v_finetune_overfit_ci",
"--num_height", "480",
"--num_width", "832",
"--num_frames", "81",
"--validation_guidance_scale", "1.0",
"--num_euler_timesteps", "50",
"--multi_phased_distill_schedule", "4000-1",
"--weight_decay", "0.01",
"--not_apply_cfg_solver",
"--dit_precision", "fp32",
"--max_grad_norm", "1.0",
]
print(f"Running training with command: {cmd}")
process = subprocess.run(cmd, check=True)
def test_e2e_overfit_single_sample():
os.environ["WANDB_MODE"] = "online"
# download_data()
run_preprocessing()
run_training()
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_v0.mp4")
print(f"reference_video_file: {reference_video_file}")
final_validation_video_file = os.path.join(LOCAL_OUTPUT_DIR, "validation_step_900_inference_steps_50_video_0.mp4")
print(f"final_validation_video_file: {final_validation_video_file}")
# Ensure both files exist
assert os.path.exists(reference_video_file), f"Reference video not found at {reference_video_file}"
assert os.path.exists(final_validation_video_file), f"Validation video not found at {final_validation_video_file}"
# Compute SSIM
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
reference_video_file,
final_validation_video_file,
use_ms_ssim=True # Using MS-SSIM for better quality assessment
)
print("\n===== SSIM Results for Step 900 Validation =====")
print(f"Mean MS-SSIM: {mean_ssim:.4f}")
print(f"Min MS-SSIM: {min_ssim:.4f}")
print(f"Max MS-SSIM: {max_ssim:.4f}")
assert max_ssim > 0.5, f"Max SSIM is below 0.5: {max_ssim}"
if __name__ == "__main__":
test_e2e_overfit_single_sample()
@@ -80,7 +80,7 @@ def run_preprocessing():
"--dataloader_num_workers", "0",
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
"--train_fps", "16",
"--validation_prompt_txt", os.path.join(LOCAL_RAW_DATA_DIR, "validation_prompt_1_sample.txt"),
"--validation_dataset_file", os.path.join(LOCAL_RAW_DATA_DIR, "validation_prompt_1_sample.json"),
"--samples_per_file", "1",
"--flush_frequency", "1",
"--video_length_tolerance_range", "5",
@@ -100,7 +100,7 @@ def run_training():
"--inference_mode", "False",
"--pretrained_model_name_or_path", MODEL_PATH,
"--data_path", LOCAL_TRAINING_DATA_DIR,
"--validation_prompt_dir", LOCAL_VALIDATION_DATA_DIR,
"--validation_preprocessed_path", LOCAL_VALIDATION_DATA_DIR,
"--train_batch_size", "1",
"--num_latent_t", "8",
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
@@ -78,7 +78,7 @@ I2V_MODEL_TO_PARAMS = {
TEST_PROMPTS = [
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.",
"A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature."
# "A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature."
]
I2V_TEST_PROMPTS = [
@@ -328,5 +328,5 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
if not success:
logger.error("Failed to write SSIM results to file")
min_acceptable_ssim = 0.97
min_acceptable_ssim = 0.95
assert mean_ssim >= min_acceptable_ssim, f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim}"
@@ -0,0 +1 @@
{"grad_norm":0.478515625,"_runtime":95.727033597,"_wandb":{"runtime":95},"_step":5,"validation_videos_50_steps":{"videos":[{"_type":"video-file","sha256":"42a1c311521a9d460db788713be1cbf2db767494e02619b43be5bf3eed8381d8","size":158632,"path":"media/videos/validation_videos_50_steps_0_42a1c311521a9d460db7.mp4"},{"path":"media/videos/validation_videos_50_steps_0_818505095b4b5e8b7f51.mp4","_type":"video-file","sha256":"818505095b4b5e8b7f511012d45f04d151ce3344bc058fc0f3225a414a851e4a","size":147825},{"sha256":"fc334ba9ed5e66c8527ee3b408e3be2d76167fef03588bf2840f4a0792f2fe34","size":136933,"path":"media/videos/validation_videos_50_steps_0_fc334ba9ed5e66c8527e.mp4","_type":"video-file"},{"size":201797,"path":"media/videos/validation_videos_50_steps_0_ccd98f6f907635d266a7.mp4","_type":"video-file","sha256":"ccd98f6f907635d266a74783688e7ecf1dac752d79d72d69eab9ef0e3f7413eb"},{"_type":"video-file","sha256":"ca79f40a0aed38f676f12779b349ce40e9e3fb7f36c578f49a20854c70508fb4","size":147114,"path":"media/videos/validation_videos_50_steps_0_ca79f40a0aed38f676f1.mp4"},{"size":175104,"path":"media/videos/validation_videos_50_steps_0_32c9b33ff920c17e5881.mp4","_type":"video-file","sha256":"32c9b33ff920c17e588133d7a27aa400ff3dc529b01ed4f16ac4d6bb2afa0f00"},{"sha256":"2cf520bfb93401c914e93c87ef791c2f12a4e043b95dfdc98115c930e11dfe67","size":139655,"path":"media/videos/validation_videos_50_steps_0_2cf520bfb93401c914e9.mp4","_type":"video-file"},{"_type":"video-file","sha256":"1d73aba17ce582c7aef4af4d64079e3e9d3df205634eff453446bdaf2340b214","size":149028,"path":"media/videos/validation_videos_50_steps_0_1d73aba17ce582c7aef4.mp4"}],"captions":false,"_type":"videos","count":8},"train_loss":0.08922439813613892,"_timestamp":1.750202051751466e+09,"avg_step_time":0.7536672964692116,"step_time":0.4742048177868128,"learning_rate":1e-05,"vsa_sparsity":0.05}
@@ -0,0 +1,136 @@
import os
import sys
import subprocess
from pathlib import Path
import json
from huggingface_hub import snapshot_download
# Import the training pipeline
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
from fastvideo.v1.training.wan_training_pipeline import main
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
wandb_name = "test_training_loss_VSA"
reference_wandb_summary_file = "fastvideo/v1/tests/training/VSA/reference_wandb_summary_VSA.json"
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "1"
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
def run_worker():
"""Worker function that will be run on each GPU"""
# Create and populate args
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
# Set the arguments as they are in finetune_v1_test.sh
args = parser.parse_args([
"--model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--inference_mode", "False",
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--cache_dir", "/home/.cache",
"--data_path", "data/mini_dataset_i2v_VSA/combined_parquet_dataset",
"--validation_preprocessed_path", "data/mini_dataset_i2v_VSA/validation_parquet_dataset",
"--train_batch_size", "1",
"--num_latent_t", "4",
"--num_gpus", "1",
"--sp_size", "1",
"--tp_size", "1",
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", "1",
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "4",
"--gradient_accumulation_steps", "1",
"--max_train_steps", "5",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--checkpointing_steps", "30",
"--validation_steps", "10",
"--validation_sampling_steps", "50",
"--log_validation",
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--cfg", "0.0",
"--output_dir", "data/wan_finetune_test_VSA",
"--tracker_project_name", "wan_finetune_ci_VSA",
"--wandb_run_name", wandb_name,
"--num_height", "384",
"--num_width", "512",
"--num_frames", "13",
"--flow_shift", "3",
"--validation_guidance_scale", "1.0",
"--num_euler_timesteps", "50",
"--weight_decay", "0.01",
"--dit_precision", "fp32",
"--max_grad_norm", "1.0",
"--VSA_decay_rate", "0.01",
"--VSA_decay_interval_steps", "1",
"--VSA_sparsity", "0.9"
])
# Call the main training function
main(args)
def test_distributed_training():
"""Test the distributed training setup"""
os.environ["WANDB_MODE"] = "online"
data_dir = Path("data/mini_dataset_i2v_VSA")
if not data_dir.exists():
print(f"Downloading test dataset to {data_dir}...")
snapshot_download(
repo_id="BrianChen1129/mini_dataset_i2v_VSA",
local_dir=str(data_dir),
repo_type="dataset",
local_dir_use_symlinks=False
)
# Get the current file path
current_file = Path(__file__).resolve()
# Run torchrun command
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE,
str(current_file)
]
process = subprocess.run(cmd, check=True)
summary_file = 'wandb/latest-run/files/wandb-summary.json'
reference_wandb_summary = json.load(open(reference_wandb_summary_file))
wandb_summary = json.load(open(summary_file))
fields_and_thresholds = {
'avg_step_time': 1.0,
'grad_norm': 0.1,
'step_time': 0.5,
'train_loss': 0.001
}
failures = []
for field, threshold in fields_and_thresholds.items():
ref_value = reference_wandb_summary[field]
current_value = wandb_summary[field]
diff = abs(ref_value - current_value)
print(f"INFO: {field}, diff: {diff}, threshold: {threshold}, reference: {ref_value}, current: {current_value}")
if diff > threshold:
failures.append(f"FAILED: {field} difference {diff} exceeds threshold of {threshold} (reference: {ref_value}, current: {current_value})")
if failures:
raise AssertionError("\n".join(failures))
if __name__ == "__main__":
if os.environ.get("LOCAL_RANK") is not None:
# We're being run by torchrun
run_worker()
else:
# We're being run directly
test_distributed_training()
@@ -18,7 +18,7 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
wandb_name = "test_training_loss"
reference_wandb_summary_file = "fastvideo/v1/tests/training/reference_wandb_summary.json"
reference_wandb_summary_file = "fastvideo/v1/tests/training/Vanilla/reference_wandb_summary.json"
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "4"
@@ -38,7 +38,7 @@ def run_worker():
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--cache_dir", "/home/.cache",
"--data_path", "data/crush-smol_parq/combined_parquet_dataset",
"--validation_prompt_dir", "data/crush-smol_parq/validation_parquet_dataset",
"--validation_preprocessed_path", "data/crush-smol_parq/validation_parquet_dataset",
"--train_batch_size", "2",
"--num_latent_t", "4",
"--num_gpus", "4",
@@ -81,7 +81,7 @@ def run_worker():
def test_distributed_training():
"""Test the distributed training setup"""
os.environ["WANDB_MODE"] = "offline"
os.environ["WANDB_MODE"] = "online"
data_dir = Path("data/crush-smol_parq")
@@ -107,14 +107,14 @@ def test_distributed_training():
process = subprocess.run(cmd, check=True)
summary_file = "fastvideo/v1/tests/training/reference_wandb_summary.json"
summary_file = 'wandb/latest-run/files/wandb-summary.json'
reference_wandb_summary = json.load(open(reference_wandb_summary_file))
wandb_summary = json.load(open(summary_file))
fields_and_thresholds = {
'avg_step_time': 1.0,
'grad_norm': 0.1,
'grad_norm': 0.2,
'step_time': 0.5,
'train_loss': 0.001
}
+478 -324
View File
@@ -2,37 +2,46 @@
import gc
import math
import os
import traceback
import time
from abc import ABC, abstractmethod
from collections import deque
from typing import Any, Dict, Iterator, List
import imageio
import numpy as np
import torch
import torch.distributed as dist
import torchvision
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.optimization import get_scheduler
from einops import rearrange
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
import fastvideo.v1.envs as envs
from fastvideo.v1.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadata)
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset import build_parquet_map_style_dataloader
from fastvideo.v1.distributed import (get_sp_group, get_torch_device,
get_world_group)
from fastvideo.v1.dataset.dataloader.schema import (
pyarrow_schema_t2v, pyarrow_schema_t2v_validation)
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
get_torch_device, get_world_group)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
TrainingBatch)
from fastvideo.v1.training.training_utils import (
compute_density_for_timestep_sampling, get_sigmas, normalize_dit_input)
clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
normalize_dit_input, save_checkpoint, shard_latents_across_sp)
from fastvideo.v1.utils import is_vsa_available, set_random_seed
import wandb # isort: skip
logger = init_logger(__name__)
vsa_available = is_vsa_available()
# Note: if checking with float32, cannot use flash-attn.
GRADIENT_CHECK_DTYPE = torch.bfloat16
logger = init_logger(__name__)
class TrainingPipeline(ComposedPipelineBase, ABC):
@@ -43,16 +52,20 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
_required_config_modules = ["scheduler", "transformer"]
validation_pipeline: ComposedPipelineBase
train_dataloader: StatefulDataLoader
train_loader_iter: Iterator[tuple[torch.Tensor, torch.Tensor, torch.Tensor,
Dict[str, Any]]]
train_loader_iter: Iterator[Dict[str, Any]]
current_epoch: int = 0
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
raise RuntimeError(
"create_pipeline_stages should not be called for training pipeline")
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_t2v
self.validation_dataset_schema = pyarrow_schema_t2v_validation
def initialize_training_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing training pipeline...")
self.training_args = training_args
self.device = get_torch_device()
world_group = get_world_group()
self.world_size = world_group.world_size
@@ -62,7 +75,12 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self.sp_world_size = self.sp_group.world_size
self.local_rank = world_group.local_rank
self.transformer = self.get_module("transformer")
assert training_args.seed is not None
self.seed = training_args.seed
assert self.transformer is not None
self.set_schemas()
# self.train_dataset_schema = pyarrow_schema_t2v
# self.validation_dataset_schema = pyarrow_schema_t2v_validation
self.transformer.requires_grad_(True)
self.transformer.train()
@@ -96,12 +114,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
training_args.data_path,
training_args.train_batch_size,
parquet_schema=self.train_dataset_schema,
num_data_workers=training_args.dataloader_num_workers,
drop_last=True,
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
text_len, # type: ignore[attr-defined]
seed=training_args.seed)
seed=self.seed)
self.noise_scheduler = noise_scheduler
@@ -130,17 +149,438 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
raise NotImplementedError(
"Training pipelines must implement this method")
@abstractmethod
def train_one_step(self, transformer, model_type, optimizer, lr_scheduler,
loader, noise_scheduler, noise_random_generator,
gradient_accumulation_steps, sp_size,
precondition_outputs, max_grad_norm, weighting_scheme,
logit_mean, logit_std, mode_scale):
"""
Train one step of the model.
"""
raise NotImplementedError(
"Training pipeline must implement this method")
def _prepare_training(self, training_batch: TrainingBatch) -> TrainingBatch:
self.transformer.requires_grad_(True)
self.transformer.train()
self.optimizer.zero_grad()
training_batch.total_loss = 0.0
return training_batch
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.train_loader_iter is not None
assert self.train_dataloader is not None
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
# latents, encoder_hidden_states, encoder_attention_mask, caption_text, extra_latents, infos = batch
# for key, value in batch.items():
# if isinstance(value, torch.Tensor):
# logger.info("key: %s, shape: %s", key, value.shape)
# else:
# logger.info("key: %s, value: %s", key, value)
# print("--------------------------------")
# logger.info("batch: %s", batch)
latents = batch['vae_latent']
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
# extra_latents = batch['extra_latents']
infos = batch['info_list']
training_batch.latents = latents.to(get_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_torch_device(), dtype=torch.bfloat16)
# training_batch.extra_latents = extra_latents
training_batch.infos = infos
return training_batch
def _normalize_dit_input(self,
training_batch: TrainingBatch) -> TrainingBatch:
# TODO(will): support other models
training_batch.latents = normalize_dit_input('wan',
training_batch.latents)
return training_batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert training_batch.latents is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert self.noise_random_generator is not None
batch_size = training_batch.latents.shape[0]
noise = torch.randn_like(training_batch.latents)
u = compute_density_for_timestep_sampling(
weighting_scheme=self.training_args.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=self.training_args.logit_mean,
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
timesteps = self.noise_scheduler.timesteps[indices].to(
device=training_batch.latents.device)
if self.training_args.sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
sp_group.broadcast(timesteps, src=0)
sigmas = get_sigmas(
self.noise_scheduler,
training_batch.latents.device,
timesteps,
n_dim=training_batch.latents.ndim,
dtype=training_batch.latents.dtype,
)
noisy_model_input = (1.0 -
sigmas) * training_batch.latents + sigmas * noise
training_batch.noisy_model_input = noisy_model_input
training_batch.timesteps = timesteps
training_batch.sigmas = sigmas
training_batch.noise = noise
return training_batch
def _build_attention_metadata(
self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
latents = training_batch.latents
assert latents is not None
assert training_batch.timesteps is not None
patch_size = self.training_args.pipeline_config.dit_config.patch_size
current_vsa_sparsity = training_batch.current_vsa_sparsity
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
dit_seq_shape = [
latents.shape[2] // patch_size[0],
latents.shape[3] // patch_size[1],
latents.shape[4] // patch_size[2]
]
training_batch.attn_metadata = VideoSparseAttentionMetadata(
current_timestep=training_batch.timesteps,
dit_seq_shape=dit_seq_shape,
VSA_sparsity=current_vsa_sparsity)
else:
training_batch.attn_metadata = None
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert training_batch.noisy_model_input is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert training_batch.timesteps is not None
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
training_batch.encoder_hidden_states,
"timestep":
training_batch.timesteps.to(get_torch_device(),
dtype=torch.bfloat16),
"encoder_attention_mask":
training_batch.encoder_attention_mask,
"return_dict":
False,
}
return training_batch
def _transformer_forward_and_compute_loss(
self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.transformer is not None
assert self.training_args is not None
assert training_batch.latents is not None
assert training_batch.noise is not None
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
assert training_batch.attn_metadata is not None
else:
assert training_batch.attn_metadata is None
assert training_batch.input_kwargs is not None
input_kwargs = training_batch.input_kwargs
# if 'hunyuan' in self.training_args.model_type:
# input_kwargs["guidance"] = torch.tensor(
# [1000.0],
# device=training_batch.noisy_model_input.device,
# dtype=torch.bfloat16)
with set_forward_context(
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
model_pred = self.transformer(**input_kwargs)
if self.training_args.precondition_outputs:
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
loss = (torch.mean((model_pred.float() - target.float())**2) /
self.training_args.gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
# local_main_process_only=False)
world_group = get_world_group()
world_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
training_batch.total_loss += avg_loss.item()
return training_batch
def _clip_grad_norm(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
max_grad_norm = self.training_args.max_grad_norm
# TODO(will): perhaps move this into transformer api so that we can do
# the following:
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
if max_grad_norm is not None:
model_parts = [self.transformer]
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
foreach=None,
)
assert grad_norm is not float('nan') or grad_norm is not float(
'inf')
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
else:
grad_norm = 0.0
training_batch.grad_norm = grad_norm
return training_batch
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
training_batch = self._prepare_training(training_batch)
for _ in range(self.training_args.gradient_accumulation_steps):
training_batch = self._get_next_batch(training_batch)
# Shard latents across sp groups
training_batch.latents = shard_latents_across_sp(
training_batch.latents,
num_latent_t=self.training_args.num_latent_t)
# Normalize DIT input
training_batch = self._normalize_dit_input(training_batch)
training_batch = self._prepare_dit_inputs(training_batch)
training_batch = self._build_attention_metadata(training_batch)
training_batch = self._build_input_kwargs(training_batch)
training_batch = self._transformer_forward_and_compute_loss(
training_batch)
training_batch = self._clip_grad_norm(training_batch)
self.optimizer.step()
self.lr_scheduler.step()
training_batch.total_loss = training_batch.total_loss
training_batch.grad_norm = training_batch.grad_norm
return training_batch
def _resume_from_checkpoint(self) -> None:
assert self.training_args is not None
logger.info("Loading checkpoint from %s",
self.training_args.resume_from_checkpoint)
resumed_step = load_checkpoint(
self.transformer, self.global_rank,
self.training_args.resume_from_checkpoint, self.optimizer,
self.train_dataloader, self.lr_scheduler,
self.noise_random_generator)
if resumed_step > 0:
self.init_steps = resumed_step
logger.info("Successfully resumed from step %s", resumed_step)
else:
logger.warning("Failed to load checkpoint, starting from step 0")
self.init_steps = 0
def train(self) -> None:
assert self.training_args is not None
# Set random seeds for deterministic training
set_random_seed(self.seed)
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
self.train_loader_iter = iter(self.train_dataloader)
step_times: deque[float] = deque(maxlen=100)
self._log_training_info()
self._log_validation(self.transformer, self.training_args, 1)
# Train!
progress_bar = tqdm(
range(0, self.training_args.max_train_steps),
initial=self.init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=self.local_rank > 0,
)
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
if vsa_available:
vsa_sparsity = self.training_args.VSA_sparsity
vsa_decay_rate = self.training_args.VSA_decay_rate
vsa_decay_interval_steps = self.training_args.VSA_decay_interval_steps
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
else:
current_vsa_sparsity = 0.0
training_batch = TrainingBatch()
training_batch.current_timestep = step
training_batch.current_vsa_sparsity = current_vsa_sparsity
training_batch = self.train_one_step(training_batch)
loss = training_batch.total_loss
grad_norm = training_batch.grad_norm
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
if self.global_rank == 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": self.lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
},
step=step,
)
if step % self.training_args.checkpointing_steps == 0:
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir, step,
self.optimizer, self.train_dataloader,
self.lr_scheduler, self.noise_random_generator)
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self._log_validation(self.transformer, self.training_args, step)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage after validation: %s MB",
gpu_memory_usage)
wandb.finish()
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir,
self.training_args.max_train_steps, self.optimizer,
self.train_dataloader, self.lr_scheduler,
self.noise_random_generator)
if get_sp_group():
cleanup_dist_env_and_memory()
def _log_training_info(self) -> None:
assert self.training_args is not None
assert self.training_args.sp_size is not None
assert self.training_args.gradient_accumulation_steps is not None
total_batch_size = (self.world_size *
self.training_args.gradient_accumulation_steps /
self.training_args.sp_size *
self.training_args.train_sp_batch_size)
logger.info("***** Running training *****")
logger.info(" Num examples = %s", len(self.train_dataset))
logger.info(" Dataloader size = %s", len(self.train_dataloader))
logger.info(" Num Epochs = %s", self.num_train_epochs)
logger.info(" Resume training from step %s",
self.init_steps) # type: ignore
logger.info(" Instantaneous batch size per device = %s",
self.training_args.train_batch_size)
logger.info(
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
total_batch_size)
logger.info(" Gradient Accumulation steps = %s",
self.training_args.gradient_accumulation_steps)
logger.info(" Total optimization steps = %s",
self.training_args.max_train_steps)
logger.info(
" Total training parameters per FSDP shard = %s B",
sum(p.numel()
for p in self.transformer.parameters() if p.requires_grad) /
1e9)
# print dtype
logger.info(" Master weight dtype: %s",
self.transformer.parameters().__next__().dtype)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage before train_one_step: %s MB",
gpu_memory_usage)
logger.info("VSA validation sparsity: %s",
self.training_args.VSA_sparsity)
def _prepare_validation_inputs(
self, sampling_param: SamplingParam, training_args: TrainingArgs,
validation_batch: Dict[str, Any], num_inference_steps: int,
negative_prompt_embeds: torch.Tensor | None,
negative_prompt_attention_mask: torch.Tensor | None
) -> ForwardBatch:
# logger.info("validation_batch: %s", validation_batch)
# latents, embeddings, masks, caption_text, extra_latents, infos = validation_batch
prompt = validation_batch['info_list'][0]['prompt']
prompt_embeds = validation_batch['text_embedding']
prompt_attention_mask = validation_batch['text_attention_mask']
prompt_embeds = prompt_embeds.to(get_torch_device())
prompt_attention_mask = prompt_attention_mask.to(get_torch_device())
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# Prepare batch for validation
batch = ForwardBatch(
prompt=prompt,
data_type="video",
latents=None,
seed=self.seed, # Use deterministic seed
generator=torch.Generator(device="cpu").manual_seed(self.seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
num_inference_steps=
num_inference_steps, # Use the current validation step
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
return batch
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
@@ -158,25 +598,25 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
# Set deterministic seed for validation
validation_seed = training_args.seed if training_args.seed is not None else 42
torch.manual_seed(validation_seed)
torch.cuda.manual_seed_all(validation_seed)
set_random_seed(self.seed)
logger.info("Using validation seed: %s", validation_seed)
logger.info("Using validation seed: %s", self.seed)
# Prepare validation prompts
logger.info('fastvideo_args.validation_prompt_dir: %s',
training_args.validation_prompt_dir)
logger.info('fastvideo_args.validation_preprocessed_path: %s',
training_args.validation_preprocessed_path)
validation_dataset, validation_dataloader = build_parquet_map_style_dataloader(
training_args.validation_prompt_dir,
training_args.validation_preprocessed_path,
batch_size=1,
parquet_schema=self.validation_dataset_schema,
num_data_workers=0,
drop_last=False,
drop_first_row=sampling_param.negative_prompt is not None,
cfg_rate=training_args.cfg)
if sampling_param.negative_prompt:
_, negative_prompt_embeds, negative_prompt_attention_mask, _ = validation_dataset.get_validation_negative_prompt(
negative_prompt_embeds, negative_prompt_attention_mask, negative_prompt = validation_dataset.get_validation_negative_prompt(
)
logger.info("negative_prompt: %s", negative_prompt)
transformer.eval()
@@ -189,43 +629,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
step_videos: List[np.ndarray] = []
step_captions: List[str | None] = []
for _, embeddings, masks, infos in validation_dataloader:
# for _, embeddings, masks, caption_text, extra_latents, infos in validation_dataloader:
for validation_batch in validation_dataloader:
batch = self._prepare_validation_inputs(
sampling_param, training_args, validation_batch,
num_inference_steps, negative_prompt_embeds,
negative_prompt_attention_mask)
step_captions.extend([None]) # TODO(peiyuan): add caption
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# Prepare batch for validation
batch = ForwardBatch(
data_type="video",
latents=None,
seed=validation_seed, # Use deterministic seed
generator=torch.Generator(
device="cpu").manual_seed(validation_seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
num_inference_steps=
num_inference_steps, # Use the current validation step
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
# Run validation inference
with torch.no_grad(), torch.autocast("cuda",
dtype=torch.bfloat16):
@@ -294,259 +704,3 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
transformer.train()
gc.collect()
torch.cuda.empty_cache()
def gradient_check_parameters(self,
transformer,
latents,
encoder_hidden_states,
encoder_attention_mask,
timesteps,
target,
eps=5e-2,
max_params_to_check=2000) -> float:
"""
Verify gradients using finite differences for FSDP models with GRADIENT_CHECK_DTYPE.
Uses standard tolerances for GRADIENT_CHECK_DTYPE precision.
"""
assert self.training_args is not None
# Move all inputs to CPU and clear GPU memory
inputs_cpu = {
'latents': latents.cpu(),
'encoder_hidden_states': encoder_hidden_states.cpu(),
'encoder_attention_mask': encoder_attention_mask.cpu(),
'timesteps': timesteps.cpu(),
'target': target.cpu()
}
del latents, encoder_hidden_states, encoder_attention_mask, timesteps, target
torch.cuda.empty_cache()
def compute_loss() -> torch.Tensor:
assert self.training_args is not None
# Move inputs to GPU, compute loss, cleanup
inputs_gpu = {
k:
v.to(get_torch_device(),
dtype=GRADIENT_CHECK_DTYPE
if k != 'encoder_attention_mask' else None)
for k, v in inputs_cpu.items()
}
# Use GRADIENT_CHECK_DTYPE for more accurate gradient checking
# with torch.autocast(enabled=False, device_type="cuda"):
with torch.autocast("cuda", dtype=GRADIENT_CHECK_DTYPE):
with set_forward_context(
current_timestep=inputs_gpu['timesteps'],
attn_metadata=None):
model_pred = transformer(
hidden_states=inputs_gpu['latents'],
encoder_hidden_states=inputs_gpu[
'encoder_hidden_states'],
timestep=inputs_gpu['timesteps'],
encoder_attention_mask=inputs_gpu[
'encoder_attention_mask'],
return_dict=False)[0]
if self.training_args.precondition_outputs:
sigmas = get_sigmas(self.noise_scheduler,
inputs_gpu['latents'].device,
inputs_gpu['timesteps'],
n_dim=inputs_gpu['latents'].ndim,
dtype=inputs_gpu['latents'].dtype)
model_pred = inputs_gpu['latents'] - model_pred * sigmas
target_adjusted = inputs_gpu['target']
else:
target_adjusted = inputs_gpu['target']
loss = torch.mean((model_pred - target_adjusted)**2)
# Cleanup and return
loss_cpu = loss.cpu()
del inputs_gpu, model_pred, target_adjusted
if 'sigmas' in locals():
del sigmas
torch.cuda.empty_cache()
return loss_cpu.to(get_torch_device())
try:
# Get analytical gradients
transformer.zero_grad()
analytical_loss = compute_loss()
analytical_loss.backward()
# Check gradients for selected parameters
absolute_errors: list[float] = []
param_count = 0
rank = dist.get_rank()
sp_group = get_sp_group()
for name, param in transformer.named_parameters():
sp_group.barrier()
# skip scale_shift_table because it is not sharded
if 'scale_shift_table' in name:
continue
if isinstance(param.grad, torch.distributed.tensor.DTensor):
full_grad = param.grad.full_tensor()
distributed = True
else:
full_grad = param.grad
distributed = False
continue
if not (param.requires_grad and param.grad is not None
and param_count < max_params_to_check
and full_grad.abs().max() > 5e-4):
continue
if not distributed and rank != 0:
continue
# Get local parameter and gradient tensors
local_param = param._local_tensor if hasattr(
param, '_local_tensor') else param
local_grad = param.grad._local_tensor if hasattr(
param.grad, '_local_tensor') else param.grad
# Find first significant gradient element
flat_param = local_param.data.view(-1)
flat_grad = local_grad.view(-1)
check_idx = next((i for i in range(min(10, flat_param.numel()))
if abs(flat_grad[i]) > 1e-4), 0)
# Store original values
orig_value = flat_param[check_idx].item()
analytical_grad = flat_grad[check_idx].item()
# Compute numerical gradient
for delta in [eps, -eps]:
with torch.no_grad():
# only have a single rank modify the parameter
# because we are using FSDP
if rank == 0:
flat_param[check_idx] = orig_value + delta
loss = compute_loss()
if delta > 0:
loss_plus = loss.item()
else:
loss_minus = loss.item()
# Restore parameter and compute error
with torch.no_grad():
flat_param[check_idx] = orig_value
numerical_grad = (loss_plus - loss_minus) / (2 * eps)
abs_error = abs(analytical_grad - numerical_grad)
rel_error = abs_error / max(abs(analytical_grad),
abs(numerical_grad), 1e-3)
absolute_errors.append(abs_error)
if self.global_rank == 0:
logger.info(
"%s[%s]: analytical=%.5f, numerical=%.5f, abs_error=%.2e, rel_error=%.2f%%",
name, check_idx, analytical_grad, numerical_grad,
abs_error, rel_error * 100)
# param_count += 1
# Compute and log statistics
if rank == 0 and absolute_errors:
min_err, max_err, mean_err = min(absolute_errors), max(
absolute_errors
), sum(absolute_errors) / len(absolute_errors)
logger.info("Gradient check stats: min=%s, max=%s, mean=%s",
min_err, max_err, mean_err)
wandb.log({
"grad_check/min_abs_error": min_err,
"grad_check/max_abs_error": max_err,
"grad_check/mean_abs_error": mean_err,
"grad_check/analytical_loss": analytical_loss.item(),
})
return max_err
return float('inf')
except Exception as e:
logger.error("Gradient check failed: %s", e)
traceback.print_exc()
return float('inf')
def setup_gradient_check(self, args, loader_iter, noise_scheduler,
noise_random_generator) -> float | None:
"""
Setup and perform gradient check on a fresh batch.
Args:
args: Training arguments
loader_iter: Data loader iterator
noise_scheduler: Noise scheduler for diffusion
noise_random_generator: Random number generator for noise
Returns:
float or None: Maximum gradient error or None if check is disabled/fails
"""
assert self.training_args is not None
try:
# Get a fresh batch and process it exactly like train_one_step
check_latents, check_encoder_hidden_states, check_encoder_attention_mask, check_infos = next(
loader_iter)
# Process exactly like in train_one_step but use GRADIENT_CHECK_DTYPE
check_latents = check_latents.to(get_torch_device(),
dtype=GRADIENT_CHECK_DTYPE)
check_encoder_hidden_states = check_encoder_hidden_states.to(
get_torch_device(), dtype=GRADIENT_CHECK_DTYPE)
check_latents = normalize_dit_input("wan", check_latents)
batch_size = check_latents.shape[0]
check_noise = torch.randn_like(check_latents)
check_u = compute_density_for_timestep_sampling(
weighting_scheme=args.weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=args.logit_mean,
logit_std=args.logit_std,
mode_scale=args.mode_scale,
)
check_indices = (check_u *
noise_scheduler.config.num_train_timesteps).long()
check_timesteps = noise_scheduler.timesteps[check_indices].to(
device=check_latents.device)
check_sigmas = get_sigmas(
noise_scheduler,
check_latents.device,
check_timesteps,
n_dim=check_latents.ndim,
dtype=check_latents.dtype,
)
check_noisy_model_input = (
1.0 - check_sigmas) * check_latents + check_sigmas * check_noise
# Compute target exactly like train_one_step
if args.precondition_outputs:
check_target = check_latents
else:
check_target = check_noise - check_latents
# Perform gradient check with the exact same inputs as training
max_grad_error = self.gradient_check_parameters(
transformer=self.transformer,
latents=
check_noisy_model_input, # Use noisy input like in training
encoder_hidden_states=check_encoder_hidden_states,
encoder_attention_mask=check_encoder_attention_mask,
timesteps=check_timesteps,
target=check_target,
max_params_to_check=100 # Check more parameters
)
if max_grad_error > 5e-2:
logger.error("❌ Large gradient error detected: %s",
max_grad_error)
else:
logger.info("✅ Gradient check passed: max error %s",
max_grad_error)
return max_grad_error
except Exception as e:
logger.error("Gradient check setup failed: %s", e)
traceback.print_exc()
return None
@@ -0,0 +1,264 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import Any, Dict
import torch
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.dataloader.schema import (
pyarrow_schema_i2v, pyarrow_schema_i2v_validation)
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
TrainingBatch)
from fastvideo.v1.pipelines.wan.wan_i2v_pipeline import (
WanImageToVideoValidationPipeline)
from fastvideo.v1.training.training_pipeline import TrainingPipeline
from fastvideo.v1.training.training_utils import shard_latents_across_sp
from fastvideo.v1.utils import is_vsa_available
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanI2VTrainingPipeline(TrainingPipeline):
"""
A training pipeline for Wan.
"""
_required_config_modules = ["scheduler", "transformer"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_i2v
self.validation_dataset_schema = pyarrow_schema_i2v_validation
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.pipeline_config.vae_config.load_encoder = False
validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
training_args.model_path,
args=None,
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)
self.validation_pipeline = validation_pipeline
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.train_loader_iter is not None
assert self.train_dataloader is not None
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
# latents, encoder_hidden_states, encoder_attention_mask, caption_text, extra_latents, infos = batch
# for key, value in batch.items():
# if isinstance(value, torch.Tensor):
# logger.info("key: %s, shape: %s", key, value.shape)
# else:
# logger.info("key: %s, value: %s", key, value)
# print("--------------------------------")
# logger.info("batch: %s", batch)
latents = batch['vae_latent']
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
clip_features = batch['clip_feature']
image_latents = batch['first_frame_latent']
pil_image = batch['pil_image']
# extra_latents = batch['extra_latents']
infos = batch['info_list']
training_batch.latents = latents.to(get_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_torch_device(), dtype=torch.bfloat16)
training_batch.preprocessed_image = pil_image.to(get_torch_device())
training_batch.image_embeds = clip_features.to(get_torch_device())
training_batch.image_latents = image_latents.to(get_torch_device())
# training_batch.extra_latents = extra_latents
training_batch.infos = infos
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert training_batch.noisy_model_input is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert training_batch.timesteps is not None
assert training_batch.preprocessed_image is not None
assert training_batch.image_embeds is not None
assert training_batch.image_latents is not None
# assert training_batch.extra_latents is not None
# extra_latents = training_batch.extra_latents
# if extra_latents:
# image_embeds, image_latents = extra_latents[
# "clip_feature"], extra_latents["first_frame_latent"]
# image_
# Image Embeds
image_embeds = training_batch.image_embeds
image_latents = training_batch.image_latents
preprocessed_image = training_batch.preprocessed_image
assert torch.isnan(image_embeds).sum() == 0
image_embeds = image_embeds.to(get_torch_device(), dtype=torch.bfloat16)
encoder_hidden_states_image = image_embeds
# Image Latents
assert torch.isnan(image_latents).sum() == 0
image_latents = image_latents.to(get_torch_device(),
dtype=torch.bfloat16)
image_latents = shard_latents_across_sp(
image_latents, num_latent_t=self.training_args.num_latent_t)
training_batch.noisy_model_input = torch.cat(
[training_batch.noisy_model_input, image_latents], dim=1)
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
training_batch.encoder_hidden_states,
"timestep":
training_batch.timesteps.to(get_torch_device(),
dtype=torch.bfloat16),
"encoder_attention_mask":
training_batch.encoder_attention_mask,
"encoder_hidden_states_image":
encoder_hidden_states_image,
"return_dict":
False,
}
return training_batch
def _prepare_validation_inputs(
self, sampling_param: SamplingParam, training_args: TrainingArgs,
validation_batch: Dict[str, Any], num_inference_steps: int,
negative_prompt_embeds: torch.Tensor | None,
negative_prompt_attention_mask: torch.Tensor | None
) -> ForwardBatch:
# latents, embeddings, masks, caption_text, extra_latents, infos = validation_batch
# latents = validation_batch['vae_latent']
embeddings = validation_batch['text_embedding']
masks = validation_batch['text_attention_mask']
clip_features = validation_batch['clip_feature']
# extra_latents = validation_batch['extra_latents']
preprocessed_image = validation_batch['pil_image']
infos = validation_batch['info_list']
prompt = infos[0]['prompt']
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
clip_features = clip_features.to(get_torch_device())
if 'colorful candies' in prompt:
logger.info("colorful candies")
from fastvideo.v1.models.vision_utils import load_video
# video_path = 'validation_dataset/yYcK4nANZz4-Scene-030.mp4'
video_path = '/mnt/user_storage/fv/FastVideo/examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-030.mp4'
video = load_video(video_path)
pil_image = video[0]
preprocessed_image = None
else:
pil_image = None
# clip_features = extra_latents.get("clip_feature")
# first_frame_latent = extra_latents.get("first_frame_latent")
# pil_image = extra_latents.get("pil_image")
# if clip_features is not None and clip_features.numel() > 0:
# clip_features = clip_features.to(get_torch_device())
# if first_frame_latent is not None and first_frame_latent.numel() > 0:
# first_frame_latent = first_frame_latent.to(get_torch_device())
# if pil_image is not None and pil_image[0] is not None and pil_image[
# 0].numel() > 0:
# pil_image = pil_image[0].to(get_torch_device())
# else:
# clip_features = None
# first_frame_latent = None
# pil_image = None
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# Prepare batch for validation
batch = ForwardBatch(
prompt=prompt,
data_type="video",
latents=None,
seed=self.seed, # Use deterministic seed
generator=torch.Generator(device="cpu").manual_seed(self.seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
image_embeds=[clip_features],
preprocessed_image=preprocessed_image,
pil_image=pil_image,
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
num_inference_steps=
num_inference_steps, # Use the current validation step
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
return batch
def main(args) -> None:
logger.info("Starting training pipeline...")
pipeline = WanI2VTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
main(args)
+3 -344
View File
@@ -1,45 +1,19 @@
# SPDX-License-Identifier: Apache-2.0
import importlib.util
import random
import sys
import time
from collections import deque
from copy import deepcopy
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from tqdm.auto import tqdm
import fastvideo.v1.envs as envs
from fastvideo.v1.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadata)
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
get_torch_device, get_world_group)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.wan.wan_pipeline import WanValidationPipeline
from fastvideo.v1.training.training_pipeline import TrainingPipeline
from fastvideo.v1.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
normalize_dit_input, save_checkpoint, shard_latents_across_sp)
from fastvideo.v1.utils import is_vsa_available
import wandb # isort: skip
vsa_available = False
if importlib.util.find_spec("vsa") is not None:
vsa_available = True
vsa_available = is_vsa_available()
logger = init_logger(__name__)
# Manual gradient checking flag - set to True to enable gradient verification
ENABLE_GRADIENT_CHECK = False
class WanTrainingPipeline(TrainingPipeline):
"""
@@ -74,321 +48,6 @@ class WanTrainingPipeline(TrainingPipeline):
self.validation_pipeline = validation_pipeline
def train_one_step( # type: ignore[override]
self,
transformer,
model_type,
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
precondition_outputs,
max_grad_norm,
weighting_scheme,
logit_mean,
logit_std,
mode_scale,
patch_size,
current_vsa_sparsity,
) -> tuple[float, float]:
assert self.training_args is not None
self.modules["transformer"].requires_grad_(True)
self.modules["transformer"].train()
total_loss = 0.0
optimizer.zero_grad()
for _ in range(gradient_accumulation_steps):
# Get next batch, handling epoch boundaries gracefully
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
latents, encoder_hidden_states, encoder_attention_mask, _ = batch
latents = latents.to(get_torch_device(), dtype=torch.bfloat16)
encoder_hidden_states = encoder_hidden_states.to(
get_torch_device(), dtype=torch.bfloat16)
latents = shard_latents_across_sp(
latents, num_latent_t=self.training_args.num_latent_t)
dit_seq_shape = [
latents.shape[2] // patch_size[0],
latents.shape[3] // patch_size[1],
latents.shape[4] // patch_size[2]
]
latents = normalize_dit_input(model_type, latents)
batch_size = latents.shape[0]
noise = torch.randn_like(latents)
u = compute_density_for_timestep_sampling(
weighting_scheme=weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=logit_mean,
logit_std=logit_std,
mode_scale=mode_scale,
)
indices = (u * noise_scheduler.config.num_train_timesteps).long()
timesteps = noise_scheduler.timesteps[indices].to(
device=latents.device)
if sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
sp_group.broadcast(timesteps, src=0)
sigmas = get_sigmas(
noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
with torch.autocast("cuda", dtype=torch.bfloat16):
input_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if 'hunyuan' in model_type:
input_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
attn_metadata = VideoSparseAttentionMetadata(
current_timestep=timesteps,
dit_seq_shape=dit_seq_shape,
VSA_sparsity=current_vsa_sparsity)
else:
attn_metadata = None
with set_forward_context(current_timestep=timesteps,
attn_metadata=attn_metadata):
model_pred = transformer(**input_kwargs)
if precondition_outputs:
model_pred = noisy_model_input - model_pred * sigmas
target = latents if precondition_outputs else noise - latents
loss = (torch.mean((model_pred.float() - target.float())**2) /
gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
# local_main_process_only=False)
world_group = get_world_group()
world_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
total_loss += avg_loss.item()
# TODO(will): perhaps move this into transformer api so that we can do
# the following:
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
if max_grad_norm is not None:
model_parts = [self.transformer]
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
foreach=None,
)
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
else:
grad_norm = 0.0
optimizer.step()
lr_scheduler.step()
return total_loss, grad_norm
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
):
assert self.training_args is not None
# Set random seeds for deterministic training
seed = self.training_args.seed if self.training_args.seed is not None else 42
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
noise_random_generator = torch.Generator(device="cpu").manual_seed(seed)
logger.info("Initialized random seeds with seed: %s", seed)
noise_scheduler = FlowMatchEulerDiscreteScheduler()
# Train!
assert self.training_args.sp_size is not None
assert self.training_args.gradient_accumulation_steps is not None
total_batch_size = (self.world_size *
self.training_args.gradient_accumulation_steps /
self.training_args.sp_size *
self.training_args.train_sp_batch_size)
logger.info("***** Running training *****")
logger.info(" Num examples = %s", len(self.train_dataset))
logger.info(" Dataloader size = %s", len(self.train_dataloader))
logger.info(" Num Epochs = %s", self.num_train_epochs)
logger.info(" Resume training from step %s",
self.init_steps) # type: ignore
logger.info(" Instantaneous batch size per device = %s",
self.training_args.train_batch_size)
logger.info(
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
total_batch_size)
logger.info(" Gradient Accumulation steps = %s",
self.training_args.gradient_accumulation_steps)
logger.info(" Total optimization steps = %s",
self.training_args.max_train_steps)
logger.info(
" Total training parameters per FSDP shard = %s B",
sum(p.numel()
for p in self.transformer.parameters() if p.requires_grad) /
1e9)
# print dtype
logger.info(" Master weight dtype: %s",
self.transformer.parameters().__next__().dtype)
if self.training_args.resume_from_checkpoint:
logger.info("Loading checkpoint from %s",
self.training_args.resume_from_checkpoint)
resumed_step = load_checkpoint(
self.transformer, self.global_rank,
self.training_args.resume_from_checkpoint, self.optimizer,
self.train_dataloader, self.lr_scheduler,
noise_random_generator)
if resumed_step > 0:
self.init_steps = resumed_step
logger.info("Successfully resumed from step %s", resumed_step)
else:
logger.warning(
"Failed to load checkpoint, starting from step 0")
self.init_steps = 0
progress_bar = tqdm(
range(0, self.training_args.max_train_steps),
initial=self.init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=self.local_rank > 0,
)
self.train_loader_iter = iter(self.train_dataloader)
step_times: deque[float] = deque(maxlen=100)
# TODO(will): fix this
# for i in range(self.init_steps):
# next(loader_iter)
# get gpu memory usage
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage before train_one_step: %s MB",
gpu_memory_usage)
logger.info("VSA validation sparsity: %s",
self.training_args.VSA_sparsity)
self._log_validation(self.transformer, self.training_args, 1)
if vsa_available:
vsa_sparsity = self.training_args.VSA_sparsity
vsa_decay_rate = self.training_args.VSA_decay_rate
vsa_decay_interval_steps = self.training_args.VSA_decay_interval_steps
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
if vsa_available:
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
else:
current_vsa_sparsity = 0.0
loss, grad_norm = self.train_one_step(
self.transformer,
# args.model_type,
"wan",
self.optimizer,
self.lr_scheduler,
self.train_loader_iter,
noise_scheduler,
noise_random_generator,
self.training_args.gradient_accumulation_steps,
self.training_args.sp_size,
self.training_args.precondition_outputs,
self.training_args.max_grad_norm,
self.training_args.weighting_scheme,
self.training_args.logit_mean,
self.training_args.logit_std,
self.training_args.mode_scale,
self.training_args.pipeline_config.dit_config.patch_size,
current_vsa_sparsity,
)
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
# Manual gradient checking - only at first step
if step == 1 and ENABLE_GRADIENT_CHECK:
logger.info("Performing gradient check at step %s", step)
self.setup_gradient_check(args, self.train_loader_iter,
noise_scheduler,
noise_random_generator)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
if self.global_rank == 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": self.lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
},
step=step,
)
if step % self.training_args.checkpointing_steps == 0:
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir, step,
self.optimizer, self.train_dataloader,
self.lr_scheduler, noise_random_generator)
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self._log_validation(self.transformer, self.training_args, step)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage after validation: %s MB",
gpu_memory_usage)
wandb.finish()
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir,
self.training_args.max_train_steps, self.optimizer,
self.train_dataloader, self.lr_scheduler,
noise_random_generator)
if get_sp_group():
cleanup_dist_env_and_memory()
def main(args) -> None:
logger.info("Starting training pipeline...")
@@ -396,7 +55,7 @@ def main(args) -> None:
pipeline = WanTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.forward(None, args)
pipeline.train()
logger.info("Training pipeline done")
+65 -1
View File
@@ -5,6 +5,7 @@ import argparse
import ctypes
import hashlib
import importlib
import importlib.util
import inspect
import json
import math
@@ -16,7 +17,7 @@ import tempfile
import threading
import traceback
from dataclasses import dataclass, fields, is_dataclass
from functools import partial, wraps
from functools import lru_cache, partial, wraps
from typing import (Any, Callable, Dict, List, Optional, Tuple, Type, TypeVar,
Union, cast)
@@ -732,3 +733,66 @@ def get_compute_dtype() -> torch.dtype:
else:
state = get_mixed_precision_state()
return state.param_dtype
def dict_to_3d_list(
mask_strategy: Optional[Dict[str, Any]] = None,
t_max: Optional[int] = None,
l_max: Optional[int] = None,
h_max: Optional[int] = None,
) -> List[List[List[Optional[torch.Tensor]]]]:
"""
Convert a dictionary of mask indices to a 3D list of tensors.
Args:
mask_strategy: keys are "t_l_h", values are torch.Tensor masks.
t_max, l_max, h_max: if provided (all three), force the output shape to (t_max, l_max, h_max).
If all three are None, infer shape from the data.
"""
# Case 1: no data, but fixed shape requested
if mask_strategy is None:
assert t_max is not None and l_max is not None and h_max is not None, (
"If mask_strategy is None, you must provide t_max, l_max, and h_max"
)
return [[[None for _ in range(h_max)] for _ in range(l_max)]
for _ in range(t_max)]
# Parse all keys into integer tuples
indices = [tuple(map(int, key.split("_"))) for key in mask_strategy]
# Decide on dimensions
if t_max is None and l_max is None and h_max is None:
# fully dynamic: infer from data
max_timesteps_idx = max(t for t, _, _ in indices) + 1
max_layer_idx = max(l for _, l, _ in indices) + 1 # noqa: E741
max_head_idx = max(h for _, _, h in indices) + 1
else:
# require all three to be provided
assert t_max is not None and l_max is not None and h_max is not None, (
"Either supply none of (t_max, l_max, h_max) to infer dimensions, "
"or supply all three to fix the shape.")
max_timesteps_idx = t_max
max_layer_idx = l_max
max_head_idx = h_max
# Preallocate
result = [[[None for _ in range(max_head_idx)]
for _ in range(max_layer_idx)] for _ in range(max_timesteps_idx)]
# Fill in, skipping any out-of-bounds entries
for key, value in mask_strategy.items():
t, l, h = map(int, key.split("_")) # noqa: E741
if 0 <= t < max_timesteps_idx and 0 <= l < max_layer_idx and 0 <= h < max_head_idx:
result[t][l][h] = value
# else: silently ignore any key that doesn’t fit
return result
def set_random_seed(seed: int) -> None:
from fastvideo.v1.platforms import current_platform
current_platform.seed_everything(seed)
@lru_cache(maxsize=1)
def is_vsa_available() -> bool:
return importlib.util.find_spec("vsa") is not None
+5 -2
View File
@@ -20,7 +20,7 @@ dependencies = [
# Machine Learning & Transformers
"transformers>=4.46.1", "tokenizers>=0.20.1", "sentencepiece==0.2.0",
"timm==1.0.11", "peft==0.13.2", "diffusers>=0.33.1", "bitsandbytes",
"torch==2.6.0", "torchvision",
"torch==2.7.1", "torchvision",
# Acceleration & Optimization
"accelerate==1.0.1",
@@ -49,6 +49,9 @@ dependencies = [
"av",
]
[tool.uv]
extra-index-url = ["https://download.pytorch.org/whl/cu128"]
[project.optional-dependencies]
# flash-attn: pip install flash-attn==2.7.4.post1 --no-cache-dir --no-build-isolation
@@ -135,4 +138,4 @@ column_limit = 80
[tool.isort]
line_length = 80
use_parentheses = true
skip_gitignore = true
skip_gitignore = true
+1 -1
View File
@@ -15,7 +15,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--data_path "$DATA_DIR"\
--validation_prompt_dir "$VALIDATION_DIR"\
--validation_preprocessed_path "$VALIDATION_DIR"\
--train_batch_size=4 \
--num_latent_t 20 \
--sp_size 4 \
+1 -1
View File
@@ -21,7 +21,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache" \
--data_path "$DATA_DIR" \
--validation_prompt_dir "$VALIDATION_DIR" \
--validation_preprocessed_path "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 16 \
--sp_size 1 \
+1 -1
View File
@@ -20,5 +20,5 @@ fastvideo generate \
--guidance-scale 5.0 \
--prompt "A beautiful woman in a red dress walking down a street" \
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--seed 12345 \
--seed 1024 \
--output-path outputs_video/

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