Compare commits

...
10 Commits
97 changed files with 1733 additions and 938 deletions
+22 -12
View File
@@ -28,7 +28,7 @@ def parse_arguments():
parser.add_argument(
'--image',
type=str,
default='runpod/pytorch:2.4.0-py3.11-cuda12.4.1-devel-ubuntu22.04',
required=True,
help='Docker image to use')
return parser.parse_args()
@@ -46,6 +46,16 @@ HEADERS = {
def create_pod():
"""Create a RunPod instance"""
# Ensure image name is lowercase (Docker requirement)
image_name = args.image.lower()
print(f"Using specified image: {image_name}")
docker_start_cmd = [
"bash",
"-c",
"apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
]
print(f"Creating RunPod instance with GPU: {args.gpu_type}...")
payload = {
"name": f"fastvideo-{JOB_ID}-{RUN_ID}",
@@ -53,8 +63,9 @@ def create_pod():
"volumeInGb": args.volume_size,
"gpuTypeIds": [args.gpu_type],
"gpuCount": args.gpu_count,
"imageName": args.image,
"allowedCudaVersions": ["12.4"]
"imageName": image_name,
"allowedCudaVersions": ["12.4"],
"dockerStartCmd": docker_start_cmd
}
response = requests.post(PODS_API, headers=HEADERS, json=payload)
@@ -91,7 +102,7 @@ def wait_for_pod(pod_id):
"Timed out waiting for RunPod to reach RUNNING state")
# Wait for ports to be assigned
max_attempts = 6
max_attempts = 50
attempts = 0
while attempts < max_attempts:
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
@@ -108,7 +119,7 @@ def wait_for_pod(pod_id):
print(
f"Waiting for SSH port and public IP to be available... (attempt {attempts+1}/{max_attempts})"
)
time.sleep(10)
time.sleep(20)
attempts += 1
if attempts >= max_attempts:
@@ -145,16 +156,15 @@ def execute_command(pod_id):
]
subprocess.run(scp_command, check=True)
# For custom image, we can use the pre-configured environment
setup_steps = [
"cd /workspace",
"wget -q https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh",
"bash Miniconda3-latest-Linux-x86_64.sh -b -p $HOME/miniconda3",
"source $HOME/miniconda3/bin/activate",
"conda create --name venv python=3.10.0 -y", "conda activate venv",
"mkdir -p /workspace/repo",
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
f"cd /workspace/{repo_name}", args.test_command
f"cd /workspace/{repo_name}",
"source /opt/conda/etc/profile.d/conda.sh",
"conda activate fastvideo-dev",
args.test_command
]
remote_command = " && ".join(setup_steps)
ssh_command = [
+78
View File
@@ -0,0 +1,78 @@
name: Build and Push Docker Image
on:
workflow_dispatch: # Only manual triggers
jobs:
build-and-push:
runs-on: ubuntu-latest
permissions:
contents: read
packages: write
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Free up disk space
run: |
# Display initial space
echo "Initial disk space:"
df -h
# Remove large directories directly
sudo rm -rf /usr/share/dotnet
sudo rm -rf /usr/local/lib/android
sudo rm -rf /opt/ghc
sudo rm -rf /usr/local/share/boost
sudo rm -rf /usr/share/swift
sudo rm -rf /usr/local/lib/node_modules
sudo rm -rf /usr/local/share/powershell
sudo rm -rf /usr/share/rust
sudo rm -rf /usr/local/.ghcup
# Remove cached files
sudo rm -rf /var/lib/apt/lists/*
sudo rm -rf /var/cache/apt/archives/*
# Clean Docker
docker system prune -af --volumes
# Display available space after cleanup
echo "Disk space after cleanup:"
df -h
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Login to GitHub Container Registry
uses: docker/login-action@v3
with:
registry: ghcr.io
username: ${{ github.repository_owner }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Extract metadata for Docker
id: meta
uses: docker/metadata-action@v5
with:
images: ghcr.io/${{ github.repository }}/fastvideo-dev
tags: |
type=raw,value=latest
type=sha,format=short
- name: Build and push Docker image
id: build-push
uses: docker/build-push-action@v6
with:
context: .
push: true
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}
cache-from: type=gha
cache-to: type=gha,mode=max
- name: Success message
run: |
echo "✅ Image successfully built and pushed to ghcr.io/${{ github.repository }}/fastvideo-dev:latest"
echo "To run tests with this image, manually trigger the 'Run Tests' workflow."
+4
View File
@@ -36,9 +36,13 @@ defaults:
shell: bash
jobs:
pre-commit:
uses: ./.github/workflows/pre-commit.yml
# Build job
build:
runs-on: ubuntu-latest
needs: pre-commit
steps:
- name: Checkout
uses: actions/checkout@v4
+2 -1
View File
@@ -6,6 +6,7 @@ on:
- main
paths:
- 'pyproject.toml' # Trigger when pyproject.toml changes
workflow_dispatch:
jobs:
check-version-change:
@@ -41,7 +42,7 @@ jobs:
build-publish-main:
needs: check-version-change
if: needs.check-version-change.outputs.version-changed == 'true'
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ubuntu-latest
permissions:
id-token: write # Needed for OIDC Trusted Publishing
+16 -12
View File
@@ -14,6 +14,11 @@ on:
- ".github/workflows/pr-test.yml"
workflow_dispatch:
inputs:
custom_image:
description: "Custom image from this repository (default: fastvideo-dev:latest)"
required: false
default: "fastvideo-dev:latest"
type: string
run_encoder_test:
description: "Run encoder-test"
required: false
@@ -35,6 +40,9 @@ on:
default: false
type: boolean
env:
PYTHONUNBUFFERED: "1"
concurrency:
group: pr-test-${{ github.ref }}
cancel-in-progress: true
@@ -107,9 +115,8 @@ jobs:
--gpu-type "NVIDIA A40"
--gpu-count 1
--volume-size 100
--test-command "pip install -e .[test] &&
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
pytest ./fastvideo/v1/tests/encoders -s"
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/encoders -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
@@ -156,9 +163,8 @@ jobs:
--gpu-type "NVIDIA A40"
--gpu-count 1
--volume-size 100
--test-command "pip install -e .[test] &&
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
pytest ./fastvideo/v1/tests/vaes -s"
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/vaes -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
@@ -205,9 +211,8 @@ jobs:
--gpu-type "NVIDIA L40S"
--gpu-count 1
--volume-size 100
--test-command "pip install -e .[test] &&
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
pytest ./fastvideo/v1/tests/transformers -s"
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/transformers -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
@@ -255,9 +260,8 @@ jobs:
--gpu-count 2
--disk-size 200
--volume-size 200
--test-command "pip install -e .[test] &&
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
pytest ./fastvideo/v1/tests/ssim -vs"
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/ssim -vs"
- name: Terminate RunPod Instances
if: ${{ always() }}
+26 -2
View File
@@ -6,6 +6,7 @@ on:
- main
paths:
- "csrc/sliding_tile_attention/setup.py"
workflow_dispatch:
jobs:
check-version-change:
@@ -43,7 +44,7 @@ jobs:
build_wheels:
name: Build Wheel
needs: check-version-change
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ${{ matrix.os }}
strategy:
@@ -57,6 +58,29 @@ jobs:
cuda-version: ['12.4.1', '12.5.1', '12.6.3']
steps:
- name: Free up disk space
run: |
echo "Initial disk space:"
df -h
# Remove large directories
sudo rm -rf /usr/share/dotnet
sudo rm -rf /usr/local/lib/android
sudo rm -rf /opt/ghc
sudo rm -rf /usr/local/share/boost
sudo rm -rf /usr/share/swift
sudo rm -rf /usr/local/lib/node_modules
sudo rm -rf /usr/local/share/powershell
sudo rm -rf /usr/share/rust
sudo rm -rf /usr/local/.ghcup
# Remove cached files
sudo rm -rf /var/lib/apt/lists/*
sudo rm -rf /var/cache/apt/archives/*
echo "Disk space after cleanup:"
df -h
- name: Checkout
uses: actions/checkout@v4
@@ -145,7 +169,7 @@ jobs:
publish_package:
name: Publish package
needs: [build_wheels, check-version-change]
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ubuntu-22.04
permissions:
id-token: write # Needed for OIDC Trusted Publishing
+29 -7
View File
@@ -3,14 +3,11 @@ __pycache__
*.pth
UCF-101/
results/
build/
fastvideo.egg-info/
wandb/
*.ipynb
*.jpg
*.safetensors
*.mp4
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
*.png
*.gif
*.pth
@@ -26,11 +23,36 @@ outputs_video
sbatch.sh
*.out
env
dist/
*.o
**/build/
**.egg-info
**.pyc
**.egg
**.txt
**.json
**.json
# Distribution / packaging
build/
dist/
*.egg-info/
*.egg
eggs/
.eggs/
# Sphinx documentation
docs/_build/
docs/source/getting_started/examples/
# VSCode
.vscode/
# DS Store
.DS_Store
# vim swap files
*.swo
*.swp
# Python pickle files
*.pkl
# Reference videos
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
+37
View File
@@ -0,0 +1,37 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
ENV DEBIAN_FRONTEND=noninteractive
WORKDIR /app
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
git \
ca-certificates \
openssh-server \
&& 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
ENV PATH=/opt/conda/bin:$PATH
RUN conda create --name fastvideo-dev python=3.10.0 -y
SHELL ["/bin/bash", "-c"]
# Copy just the pyproject.toml first to leverage Docker cache
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.0.post2 --no-build-isolation && \
conda clean -afy
COPY . .
EXPOSE 22
+7 -1
View File
@@ -5,7 +5,7 @@
FastVideo is a lightweight framework for accelerating large video diffusion models.
<p align="center">
🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank">FastHunyuan</a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank">FastMochi</a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> Slack </a>
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</b> </a> |
</p>
https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1
@@ -44,6 +44,12 @@ pip install flash-attn==2.7.0.post2
To try Sliding Tile Attention (optional), please follow the instruction in [csrc/sliding_tile_attention/README.md](csrc/sliding_tile_attention/README.md) to install STA.
You can also install the Sliding Tile Attention package using
```
pip install st_attn==0.0.3
```
## 🚀 Inference
### Inference StepVideo with Sliding Tile Attention
First, download the model:
+1 -1
View File
@@ -9,7 +9,7 @@ target = target.lower()
# Package metadata
PACKAGE_NAME = "st_attn"
VERSION = "0.0.2"
VERSION = "0.0.3"
AUTHOR = "Hao AI Lab"
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
@@ -446,7 +446,7 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
auto threads = NUM_WORKERS * kittens::WARP_THREADS;
if (has_text) {
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4-2), qo_heads, batch);
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
dim3 grid_text(2, qo_heads, batch);
if (!process_text) {
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
+3 -1
View File
@@ -1,3 +1,5 @@
(developer-guide)=
# Contributing to FastVideo
Thank you for your interest in contributing to FastVideo. We want to make the process as smooth for you as possible and this is a guide to help get you started!
@@ -27,7 +29,7 @@ conda activate fastvideo
Clone the FastVideo repository and go to the FastVideo directory:
```
git clone https://github.com/vllm-project/vllm.git && cd vllm
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
```
+21 -20
View File
@@ -1,4 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
import itertools
import re
@@ -177,28 +178,28 @@ def generate_examples():
# Category indices stored in reverse order because they are inserted into
# examples_index.documents at index 0 in order
category_indices = {
"other":
# "other":
# Index(
# path=EXAMPLE_DOC_DIR / "examples_other_index.md",
# title="Other",
# description=
# "Other examples that don't strongly fit into the online or offline serving categories.", # noqa: E501
# caption="Examples",
# ),
# "online_serving":
# Index(
# path=EXAMPLE_DOC_DIR / "examples_online_serving_index.md",
# title="Online Serving",
# description=
# "Online serving examples demonstrate how to use FastVideo in an online setting, where the model is queried for predictions in real-time.", # noqa: E501
# caption="Examples",
# ),
"inference":
Index(
path=EXAMPLE_DOC_DIR / "examples_other_index.md",
title="Other",
path=EXAMPLE_DOC_DIR / "examples_inference_index.md",
title="Inference",
description=
"Other examples that don't strongly fit into the online or offline serving categories.", # noqa: E501
caption="Examples",
),
"online_serving":
Index(
path=EXAMPLE_DOC_DIR / "examples_online_serving_index.md",
title="Online Serving",
description=
"Online serving examples demonstrate how to use FastVideo in an online setting, where the model is queried for predictions in real-time.", # noqa: E501
caption="Examples",
),
"offline_inference":
Index(
path=EXAMPLE_DOC_DIR / "examples_offline_inference_index.md",
title="Offline Inference",
description=
"Offline inference examples demonstrate how to use FastVideo in an offline setting, where the model is queried for predictions in batches. We recommend starting with <project:basic.md>.", # noqa: E501
"Inference examples demonstrate how to use FastVideo in an offline setting, where the model is queried for predictions in batches. We recommend starting with <project:basic.md>.", # noqa: E501
caption="Examples",
),
}
@@ -1,10 +0,0 @@
# Examples
A collection of examples demonstrating usage of FastVideo.
All documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.
:::{toctree}
:caption: Examples
:maxdepth: 2
:::
+79 -4
View File
@@ -1,10 +1,85 @@
(fastvideo-installation)=
# 🔧 Installation
The code is tested on Python 3.10.0, CUDA 12.4 and H100.
```
./env_setup.sh fastvideo
FastVideo currently only supports Linux and CUDA GPUs. The code is tested on Python 3.10.0 and CUDA 12.4, primarily with NVIDIA H100 GPUs.
## Prerequisites
- CUDA 12.4 installed and supported
- Linux operating system
## Installation Options
### Option 1: Quick Install
```bash
pip install fastvideo
```
To try Sliding Tile Attention (optional), please follow the instruction in [here](#sta-installation) to install STA.
### Option 2: Installation from Source
#### 1. Install Miniconda (if not already installed)
```bash
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
#### 2. Create and activate a Conda environment for FastVideo
```bash
conda create -n fastvideo python=3.10 -y
conda activate fastvideo
```
#### 3. Clone the FastVideo repository
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
```
#### 4. Install FastVideo
Basic installation:
```bash
pip install -e .
```
## Optional Dependencies
### Flash Attention
```bash
pip install flash-attn==2.7.0.post2 --no-build-isolation
```
### Sliding Tile Attention (STA)
To try Sliding Tile Attention (optional), please follow the instructions in [csrc/sliding_tile_attention/README.md](#sta-installation) to install STA.
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
[Contributor Guide](#developer-guide)
## Hardware Requirements
### For Basic Inference
- NVIDIA GPU with CUDA support
- Minimum 20GB VRAM for quantized models (e.g., single RTX 4090)
### For Lora Finetuning
- 40GB GPU memory each for 2 GPUs with lora
- 30GB GPU memory each for 2 GPUs with CPU offload and lora
### For Full Finetuning/Distillation
- Multiple high-memory GPUs recommended (e.g., H100)
## Troubleshooting
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg) for additional support.
+3
View File
@@ -0,0 +1,3 @@
# Basic
The class provides the main python interface for using FastVideo's inference pipeline.
+1
View File
@@ -0,0 +1 @@
print('Hello, world!')
+1
View File
@@ -0,0 +1 @@
View File
+4 -3
View File
@@ -7,7 +7,7 @@ from typing import (TYPE_CHECKING, Any, Dict, Generic, Optional, Protocol, Set,
Type, TypeVar)
if TYPE_CHECKING:
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
import torch
@@ -154,7 +154,7 @@ class AttentionMetadataBuilder(ABC, Generic[T]):
self,
current_timestep: int,
forward_batch: "ForwardBatch",
inference_args: "InferenceArgs",
fastvideo_args: "FastVideoArgs",
) -> T:
"""Build attention metadata with on-device tensors."""
raise NotImplementedError
@@ -186,9 +186,10 @@ class AttentionImpl(ABC, Generic[T]):
num_heads: int,
head_size: int,
softmax_scale: float,
dropout_rate: float = 0.0,
causal: bool = False,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
raise NotImplementedError
+18 -9
View File
@@ -3,7 +3,16 @@
from typing import List, Optional, Type
import torch
from flash_attn import flash_attn_func
from flash_attn import flash_attn_func as flash_attn_2_func
try:
from flash_attn_interface import flash_attn_func as flash_attn_3_func
# flash_attn 3 has slightly different API: it returns lse by default
flash_attn_func = lambda q, k, v, softmax_scale, causal: flash_attn_3_func(
q, k, v, softmax_scale, causal)[0]
except ImportError:
flash_attn_func = flash_attn_2_func
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
@@ -45,12 +54,12 @@ class FlashAttentionImpl(AttentionImpl):
self,
num_heads: int,
head_size: int,
dropout_rate: float,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.dropout_rate = dropout_rate
self.causal = causal
self.softmax_scale = softmax_scale
@@ -61,10 +70,10 @@ class FlashAttentionImpl(AttentionImpl):
value: torch.Tensor,
attn_metadata: AttentionMetadata,
):
output = flash_attn_func(query,
key,
value,
dropout_p=self.dropout_rate,
softmax_scale=self.softmax_scale,
causal=self.causal)
output = flash_attn_func(
query, # type: ignore[no-untyped-call]
key,
value,
softmax_scale=self.softmax_scale,
causal=self.causal)
return output
+4 -3
View File
@@ -38,14 +38,15 @@ class SDPAImpl(AttentionImpl):
self,
num_heads: int,
head_size: int,
dropout_rate: float,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.dropout_rate = dropout_rate
self.causal = causal
self.softmax_scale = softmax_scale
self.dropout = extra_impl_args.get("dropout_p", 0.0)
def forward(
self,
@@ -60,7 +61,7 @@ class SDPAImpl(AttentionImpl):
value = value.transpose(1, 2)
attn_kwargs = {
"attn_mask": None,
"dropout_p": self.dropout_rate,
"dropout_p": self.dropout,
"is_causal": self.causal,
"scale": self.softmax_scale
}
@@ -12,7 +12,7 @@ from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.v1.distributed import get_sp_group
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -62,7 +62,7 @@ class SlidingTileAttentionBackend(AttentionBackend):
@dataclass
class SlidingTileAttentionMetadata(AttentionMetadata):
text_length: int
current_timestep: int
class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
@@ -77,13 +77,10 @@ class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
self,
current_timestep: int,
forward_batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> SlidingTileAttentionMetadata:
return SlidingTileAttentionMetadata(
current_timestep=current_timestep,
text_length=forward_batch.attention_mask.sum(),
)
return SlidingTileAttentionMetadata(current_timestep=current_timestep, )
class SlidingTileAttentionImpl(AttentionImpl):
@@ -92,10 +89,11 @@ class SlidingTileAttentionImpl(AttentionImpl):
self,
num_heads: int,
head_size: int,
dropout_rate: float,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
# TODO(will-refactor): for now this is the mask strategy, but maybe we should
# have a more general config for STA?
@@ -107,7 +105,7 @@ class SlidingTileAttentionImpl(AttentionImpl):
mask_strategy = json.load(f)
mask_strategy = dict_to_3d_list(mask_strategy)
self.prefix = prefix
self.mask_strategy = mask_strategy
sp_group = get_sp_group()
self.sp_size = sp_group.world_size
@@ -172,8 +170,11 @@ class SlidingTileAttentionImpl(AttentionImpl):
assert self.mask_strategy[
0] is not None, "mask_strategy[0] cannot be None for SlidingTileAttention"
text_length = attn_metadata.text_length
timestep = attn_metadata.current_timestep
# pattern:'.double_blocks.0.attn.impl' or '.single_blocks.0.attn.impl'
layer_idx = int(self.prefix.split('.')[-3])
# TODO: remove hardcode
text_length = q.shape[1] - (30 * 48 * 80)
query = q.transpose(1, 2)
key = k.transpose(1, 2)
value = v.transpose(1, 2)
@@ -183,13 +184,10 @@ class SlidingTileAttentionImpl(AttentionImpl):
current_rank = sp_group.rank_in_group
start_head = current_rank * head_num
windows = [
self.mask_strategy[head_idx + start_head]
self.mask_strategy[timestep][layer_idx][head_idx + start_head]
for head_idx in range(head_num)
]
hidden_states = sliding_tile_attention(query, key, value, windows,
text_length).transpose(1, 2)
hidden_states = hidden_states.transpose(1, 2)
return hidden_states
+15 -14
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional
from typing import List, Optional
import torch
import torch.nn as nn
@@ -12,6 +12,7 @@ from fastvideo.v1.distributed.communication_op import (
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size)
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.platforms import _Backend
class DistributedAttention(nn.Module):
@@ -22,13 +23,12 @@ class DistributedAttention(nn.Module):
num_heads: int,
head_size: int,
num_kv_heads: Optional[int] = None,
dropout_rate: float = 0.0,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[List[_Backend]] = None,
prefix: str = "",
**extra_impl_args) -> None:
super().__init__()
# self.dropout_rate = dropout_rate
# self.causal = causal
if softmax_scale is None:
self.softmax_scale = head_size**-0.5
else:
@@ -38,14 +38,17 @@ class DistributedAttention(nn.Module):
num_kv_heads = num_heads
dtype = torch.get_default_dtype()
attn_backend = get_attn_backend(head_size, dtype, distributed=True)
attn_backend = get_attn_backend(
head_size,
dtype,
supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.impl = impl_cls(num_heads=num_heads,
head_size=head_size,
dropout_rate=dropout_rate,
causal=causal,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
prefix=f"{prefix}.impl",
**extra_impl_args)
self.num_heads = num_heads
self.head_size = head_size
@@ -97,7 +100,6 @@ class DistributedAttention(nn.Module):
qkv = sequence_model_parallel_all_to_all_4D(qkv,
scatter_dim=2,
gather_dim=1)
# Apply backend-specific preprocess_qkv
qkv = self.impl.preprocess_qkv(qkv, ctx_attn_metadata)
@@ -124,8 +126,7 @@ class DistributedAttention(nn.Module):
output = output[:, :seq_len * world_size]
# TODO: make this asynchronous
replicated_output = sequence_model_parallel_all_gather(
replicated_output, dim=2)
replicated_output.contiguous(), dim=2)
# Apply backend-specific postprocess_output
output = self.impl.postprocess_output(output, ctx_attn_metadata)
@@ -143,13 +144,11 @@ class LocalAttention(nn.Module):
num_heads: int,
head_size: int,
num_kv_heads: Optional[int] = None,
dropout_rate: float = 0.0,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[List[_Backend]] = None,
**extra_impl_args) -> None:
super().__init__()
# self.dropout_rate = dropout_rate
# self.causal = causal
if softmax_scale is None:
self.softmax_scale = head_size**-0.5
else:
@@ -158,11 +157,13 @@ class LocalAttention(nn.Module):
num_kv_heads = num_heads
dtype = torch.get_default_dtype()
attn_backend = get_attn_backend(head_size, dtype, distributed=False)
attn_backend = get_attn_backend(
head_size,
dtype,
supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.impl = impl_cls(num_heads=num_heads,
head_size=head_size,
dropout_rate=dropout_rate,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
causal=causal,
+7 -20
View File
@@ -3,8 +3,7 @@
import os
from contextlib import contextmanager
from functools import cache
from typing import Generator, Optional, Type, cast
from typing import Generator, List, Optional, Type, cast
import torch
@@ -82,29 +81,15 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_attn_backend(
head_size: int,
dtype: torch.dtype,
distributed: bool,
) -> Type[AttentionBackend]:
"""Selects which attention backend to use and lazily imports it."""
# Accessing envs.* behind an @lru_cache decorator can cause the wrong
# value to be returned from the cache if the value changes between calls.
return _cached_get_attn_backend(
head_size=head_size,
dtype=dtype,
distributed=distributed,
)
@cache
def _cached_get_attn_backend(
head_size: int,
dtype: torch.dtype,
distributed: bool,
supported_attention_backends: Optional[List[_Backend]] = None,
) -> Type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
#
# THIS SELECTION OVERRIDES THE FASTVIDEO_ATTENTION_BACKEND
# ENVIRONMENT VARIABLE.
if not supported_attention_backends:
raise ValueError("supported_attention_backends is empty")
selected_backend = None
backend_by_global_setting: Optional[_Backend] = (
get_global_forced_attn_backend())
@@ -117,8 +102,10 @@ def _cached_get_attn_backend(
selected_backend = backend_name_to_enum(backend_by_env_var)
# get device-specific attn_backend
if selected_backend not in supported_attention_backends:
selected_backend = None
attention_cls = current_platform.get_attn_backend_cls(
selected_backend, head_size, dtype, distributed)
selected_backend, head_size, dtype)
if not attention_cls:
raise ValueError(
f"Invalid attention backend for {current_platform.device_name}")
+8
View File
@@ -0,0 +1,8 @@
from fastvideo.v1.configs.hunyuan import HunyuanConfig, FastHunyuanConfig
from fastvideo.v1.configs.wan import WanT2V480PConfig, WanI2V480PConfig
from fastvideo.v1.configs.base import BaseConfig, SlidingTileAttnConfig
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig"
]
+70
View File
@@ -0,0 +1,70 @@
from dataclasses import dataclass, field
from typing import Optional, Dict, Any
@dataclass
class BaseConfig:
"""Base configuration for all pipeline architectures."""
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 125
fps: int = 24
# Video generation parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
seed: int = 1024
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
vae_scale_factor: Optional[int] = None
# DiT configuration
num_channels_latents: Optional[int] = None
# Image encoder configuration
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_len: int = -1
hidden_state_skip_layer: int = 0
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
neg_prompt: Optional[str] = None
# Additional parameters can be added as a dict
extra_params: Dict[str, Any] = field(default_factory=dict)
@dataclass
class SlidingTileAttnConfig(BaseConfig):
"""Configuration for sliding tile attention."""
# Override any BaseConfig defaults as needed
# Add sliding tile specific parameters
window_size: int = 16
stride: int = 8
# You can provide custom defaults for inherited fields
height: int = 576
width: int = 1024
# Additional configuration specific to sliding tile attention
pad_to_square: bool = False
use_overlap_optimization: bool = True
+39
View File
@@ -0,0 +1,39 @@
from dataclasses import dataclass
from fastvideo.v1.configs.base import BaseConfig
@dataclass
class HunyuanConfig(BaseConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
# Denoising stage
embedded_cfg_scale: int = 6
flow_shift: int = 7
num_inference_steps: int = 50
# Text encoding stage
hidden_state_skip_layer: int = 2
text_len: int = 256
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precision: str = "fp16"
# HunyuanConfig-specific added parameters
# Secondary text encoder
text_encoder_precision_2: str = "fp16"
text_len_2: int = 77
@dataclass
class FastHunyuanConfig(HunyuanConfig):
"""Configuration specifically optimized for FastHunyuan weights."""
# Override HunyuanConfig defaults
num_inference_steps: int = 6
flow_shift: int = 17
# No need to re-specify guidance_scale or embedded_cfg_scale as they
# already have the desired values from HunyuanConfig
+77
View File
@@ -0,0 +1,77 @@
"""Registry for pipeline weight-specific configurations."""
import os
from typing import Dict, Type, Optional, Callable
from fastvideo.v1.configs.base import BaseConfig
from fastvideo.v1.configs.hunyuan import HunyuanConfig, FastHunyuanConfig
from fastvideo.v1.configs.wan import WanT2V480PConfig, WanI2V480PConfig
from fastvideo.v1.utils import maybe_download_model_index, verify_model_config_and_directory
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[BaseConfig]] = {
"FastVideo/FastHunyuan-Diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig
# Add other specific weight variants
}
# For determining pipeline type from model ID
PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
"hunyuan": lambda id: "hunyuan" in id.lower(),
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
# Add other pipeline architecture detectors
}
# Fallback configs when exact match isn't found but architecture is detected
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[BaseConfig]] = {
"hunyuan":
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
"wanpipeline":
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
# Other fallbacks by architecture
}
def get_pipeline_config_for_name(
pipeline_name_or_path: str) -> Optional[Type[BaseConfig]]:
"""Get the appropriate config class for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
logger.warning(
"FastVideo may not correctly identify the optimal config for this model, as the local directory may have been renamed."
)
else:
config = maybe_download_model_index(pipeline_name_or_path)
pipeline_name = config["_class_name"]
# First try exact match for specific weights
if pipeline_name_or_path in WEIGHT_CONFIG_REGISTRY:
return WEIGHT_CONFIG_REGISTRY[pipeline_name_or_path]
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in WEIGHT_CONFIG_REGISTRY.items():
if registered_id in pipeline_name_or_path:
return config_class
# If no match, try to use the fallback config
fallback_config = None
print(pipeline_name)
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in PIPELINE_DETECTOR.items():
if detector(pipeline_name.lower()):
fallback_config = PIPELINE_FALLBACK_CONFIG.get(pipeline_type)
break
logger.warning("No match found for pipeline %s, using fallback config %s.",
pipeline_name_or_path, fallback_config)
return fallback_config
+44
View File
@@ -0,0 +1,44 @@
from dataclasses import dataclass
from fastvideo.v1.configs.base import BaseConfig
@dataclass
class WanT2V480PConfig(BaseConfig):
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
# WanConfig-specific parameters with defaults
# Video parameters
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
use_cpu_offload: bool = True
# Denoising stage
guidance_scale: float = 3.0
neg_prompt: str = "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"
flow_shift: int = 3
num_inference_steps: int = 50
# Text encoding stage
text_len: int = 512
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precision: str = "fp32"
# WanConfig-specific added parameters
@dataclass
class WanI2V480PConfig(WanT2V480PConfig):
"""Base configuration for Wan I2V 14B 480P pipeline architecture."""
# WanConfig-specific parameters with defaults
# Denoising stage
guidance_scale: float = 5.0
num_inference_steps: int = 40
# Precision for each component
image_encoder_precision: str = "fp32"
@@ -190,5 +190,5 @@ class DeviceCommunicatorBase:
torch.distributed.recv(tensor, self.ranks[src], self.device_group)
return tensor
def destroy(self):
def destroy(self) -> None:
pass
+2 -2
View File
@@ -845,12 +845,12 @@ def initialize_model_parallel(
group_name="sp")
def get_sequence_model_parallel_world_size():
def get_sequence_model_parallel_world_size() -> int:
"""Return world size for the sequence model parallel group."""
return get_sp_group().world_size
def get_sequence_model_parallel_rank():
def get_sequence_model_parallel_rank() -> int:
"""Return my rank for the sequence model parallel group."""
return get_sp_group().rank_in_group
+1 -1
View File
@@ -9,7 +9,7 @@ from fastvideo.v1.utils import FlexibleArgumentParser
class CLISubcommand:
"""Base class for CLI subcommands"""
def __init__(self):
def __init__(self) -> None:
self.name = ""
def cmd(self, args: argparse.Namespace) -> None:
+4 -4
View File
@@ -2,11 +2,11 @@
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
import argparse
from typing import List
from typing import List, cast
from fastvideo.v1.entrypoints.cli import utils
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.utils import FlexibleArgumentParser
@@ -82,9 +82,9 @@ class GenerateSubcommand(CLISubcommand):
default=None,
help="Port for the master process")
generate_parser = InferenceArgs.add_cli_args(generate_parser)
generate_parser = FastVideoArgs.add_cli_args(generate_parser)
return generate_parser
return cast(FlexibleArgumentParser, generate_parser)
def cmd_init() -> List[CLISubcommand]:
+4 -1
View File
@@ -3,13 +3,16 @@
import os
import subprocess
import sys
from typing import List, Optional
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def launch_distributed(num_gpus=None, args=None, master_port=None):
def launch_distributed(num_gpus: int,
args: List[str],
master_port: Optional[int] = None) -> int:
"""
Launch a distributed job with the given arguments
-14
View File
@@ -8,8 +8,6 @@ if TYPE_CHECKING:
FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60
FASTVIDEO_NCCL_SO_PATH: Optional[str] = None
LD_LIBRARY_PATH: Optional[str] = None
FASTVIDEO_USE_TRITON_FLASH_ATTN: bool = False
FASTVIDEO_FLASH_ATTN_VERSION: Optional[int] = None
LOCAL_RANK: int = 0
CUDA_VISIBLE_DEVICES: Optional[str] = None
FASTVIDEO_CACHE_ROOT: str = os.path.expanduser("~/.cache/fastvideo")
@@ -127,18 +125,6 @@ environment_variables: Dict[str, Callable[[], Any]] = {
"LD_LIBRARY_PATH":
lambda: os.environ.get("LD_LIBRARY_PATH", None),
# flag to control if fastvideo should use triton flash attention
"FASTVIDEO_USE_TRITON_FLASH_ATTN":
lambda:
(os.environ.get("FASTVIDEO_USE_TRITON_FLASH_ATTN", "True").lower() in
("true", "1")),
# Force fastvideo to use a specific flash-attention version (2 or 3), only valid
# when using the flash-attention backend.
"FASTVIDEO_FLASH_ATTN_VERSION":
lambda: maybe_convert_int(
os.environ.get("FASTVIDEO_FLASH_ATTN_VERSION", None)),
# Internal flag to enable Dynamo fullgraph capture
"FASTVIDEO_TEST_DYNAMO_FULLGRAPH_CAPTURE":
lambda: bool(
@@ -4,23 +4,33 @@
import argparse
import dataclasses
from contextlib import contextmanager
from typing import List, Optional
from fastvideo.v1.utils import FlexibleArgumentParser
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@dataclasses.dataclass
class InferenceArgs:
class FastVideoArgs:
# Model and path configuration
model_path: str
# Distributed executor backend
distributed_executor_backend: str = "torch"
inference_mode: bool = True # if False == training mode
# HuggingFace specific parameters
trust_remote_code: bool = False
revision: Optional[str] = None
# Parallelism
tp_size: int = 1
sp_size: int = 1
num_gpus: int = 1
tp_size: Optional[int] = None
sp_size: Optional[int] = None
dist_timeout: Optional[int] = None # timeout for torch.distributed
# Video generation parameters
@@ -42,6 +52,10 @@ class InferenceArgs:
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = False
vae_scale_factor: Optional[int] = None
# DiT configuration
num_channels_latents: Optional[int] = None
# Image encoder configuration
image_encoder_precision: str = "fp32"
@@ -108,40 +122,55 @@ class InferenceArgs:
help="Directory containing StepVideo model",
)
# distributed_executor_backend
parser.add_argument(
"--distributed-executor-backend",
type=str,
choices=["mp", "ray", "torch"],
default=FastVideoArgs.distributed_executor_backend,
help="The distributed executor backend to use",
)
# HuggingFace specific parameters
parser.add_argument(
"--trust-remote-code",
action="store_true",
default=InferenceArgs.trust_remote_code,
default=FastVideoArgs.trust_remote_code,
help="Trust remote code when loading HuggingFace models",
)
parser.add_argument(
"--revision",
type=str,
default=InferenceArgs.revision,
default=FastVideoArgs.revision,
help=
"The specific model version to use (can be a branch name, tag name, or commit id)",
)
# Parallelism
parser.add_argument(
"--num-gpus",
type=int,
default=FastVideoArgs.num_gpus,
help="The number of GPUs to use.",
)
parser.add_argument(
"--tensor-parallel-size",
"--tp-size",
type=int,
default=InferenceArgs.tp_size,
default=FastVideoArgs.tp_size,
help="The tensor parallelism size.",
)
parser.add_argument(
"--sequence-parallel-size",
"--sp-size",
type=int,
default=InferenceArgs.sp_size,
default=FastVideoArgs.sp_size,
help="The sequence parallelism size.",
)
parser.add_argument(
"--dist-timeout",
type=int,
default=InferenceArgs.dist_timeout,
default=FastVideoArgs.dist_timeout,
help="Set timeout for torch.distributed initialization.",
)
@@ -149,56 +178,56 @@ class InferenceArgs:
parser.add_argument(
"--height",
type=int,
default=InferenceArgs.height,
default=FastVideoArgs.height,
help="Height of generated video",
)
parser.add_argument(
"--width",
type=int,
default=InferenceArgs.width,
default=FastVideoArgs.width,
help="Width of generated video",
)
parser.add_argument(
"--num-frames",
type=int,
default=InferenceArgs.num_frames,
default=FastVideoArgs.num_frames,
help="Number of frames to generate",
)
parser.add_argument(
"--num-inference-steps",
type=int,
default=InferenceArgs.num_inference_steps,
default=FastVideoArgs.num_inference_steps,
help="Number of inference steps",
)
parser.add_argument(
"--guidance-scale",
type=float,
default=InferenceArgs.guidance_scale,
default=FastVideoArgs.guidance_scale,
help="Guidance scale for classifier-free guidance",
)
parser.add_argument(
"--guidance-rescale",
type=float,
default=InferenceArgs.guidance_rescale,
default=FastVideoArgs.guidance_rescale,
help="Guidance rescale for classifier-free guidance",
)
parser.add_argument(
"--embedded-cfg-scale",
type=float,
default=InferenceArgs.embedded_cfg_scale,
default=FastVideoArgs.embedded_cfg_scale,
help="Embedded CFG scale",
)
parser.add_argument(
"--flow-shift",
"--shift",
type=float,
default=InferenceArgs.flow_shift,
default=FastVideoArgs.flow_shift,
help="Flow shift parameter",
)
parser.add_argument(
"--output-type",
type=str,
default=InferenceArgs.output_type,
default=FastVideoArgs.output_type,
choices=["pil"],
help="Output type for the generated video",
)
@@ -206,7 +235,7 @@ class InferenceArgs:
parser.add_argument(
"--precision",
type=str,
default=InferenceArgs.precision,
default=FastVideoArgs.precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for the model",
)
@@ -215,14 +244,14 @@ class InferenceArgs:
parser.add_argument(
"--vae-precision",
type=str,
default=InferenceArgs.vae_precision,
default=FastVideoArgs.vae_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for VAE",
)
parser.add_argument(
"--vae-tiling",
action="store_true",
default=InferenceArgs.vae_tiling,
default=FastVideoArgs.vae_tiling,
help="Enable VAE tiling",
)
parser.add_argument(
@@ -234,14 +263,14 @@ class InferenceArgs:
parser.add_argument(
"--text-encoder-precision",
type=str,
default=InferenceArgs.text_encoder_precision,
default=FastVideoArgs.text_encoder_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for text encoder",
)
parser.add_argument(
"--text-len",
type=int,
default=InferenceArgs.text_len,
default=FastVideoArgs.text_len,
help="Maximum text length",
)
@@ -249,7 +278,7 @@ class InferenceArgs:
parser.add_argument(
"--image-encoder-precision",
type=str,
default=InferenceArgs.image_encoder_precision,
default=FastVideoArgs.image_encoder_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for image encoder",
)
@@ -259,14 +288,14 @@ class InferenceArgs:
parser.add_argument(
"--text-encoder-precision-2",
type=str,
default=InferenceArgs.text_encoder_precision_2,
default=FastVideoArgs.text_encoder_precision_2,
choices=["fp32", "fp16", "bf16"],
help="Precision for secondary text encoder",
)
parser.add_argument(
"--text-len-2",
type=int,
default=InferenceArgs.text_len_2,
default=FastVideoArgs.text_len_2,
help="Maximum secondary text length",
)
@@ -274,13 +303,13 @@ class InferenceArgs:
parser.add_argument(
"--flow-solver",
type=str,
default=InferenceArgs.flow_solver,
default=FastVideoArgs.flow_solver,
help="Solver for flow matching",
)
parser.add_argument(
"--denoise-type",
type=str,
default=InferenceArgs.denoise_type,
default=FastVideoArgs.denoise_type,
help="Denoise type for noised inputs",
)
@@ -301,7 +330,7 @@ class InferenceArgs:
parser.add_argument(
"--scheduler-type",
type=str,
default=InferenceArgs.scheduler_type,
default=FastVideoArgs.scheduler_type,
help="Type of scheduler to use",
)
@@ -309,19 +338,19 @@ class InferenceArgs:
parser.add_argument(
"--neg-prompt",
type=str,
default=InferenceArgs.neg_prompt,
default=FastVideoArgs.neg_prompt,
help="Negative prompt for sampling",
)
parser.add_argument(
"--num-videos",
type=int,
default=InferenceArgs.num_videos,
default=FastVideoArgs.num_videos,
help="Number of videos to generate per prompt",
)
parser.add_argument(
"--fps",
type=int,
default=InferenceArgs.fps,
default=FastVideoArgs.fps,
help="Frames per second for output video",
)
parser.add_argument(
@@ -340,7 +369,7 @@ class InferenceArgs:
parser.add_argument(
"--log-level",
type=str,
default=InferenceArgs.log_level,
default=FastVideoArgs.log_level,
help="The logging level of all loggers.",
)
@@ -364,20 +393,20 @@ class InferenceArgs:
parser.add_argument(
"--output-path",
type=str,
default=InferenceArgs.output_path,
default=FastVideoArgs.output_path,
help="Directory to save generated videos",
)
parser.add_argument(
"--seed",
type=int,
default=InferenceArgs.seed,
default=FastVideoArgs.seed,
help="Random seed for reproducibility",
)
return parser
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "InferenceArgs":
def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs":
args.tp_size = args.tensor_parallel_size
args.sp_size = args.sequence_parallel_size
args.flow_shift = getattr(args, "shift", args.flow_shift)
@@ -404,6 +433,15 @@ class InferenceArgs:
def check_inference_args(self) -> None:
"""Validate inference arguments for consistency"""
if self.tp_size is None:
self.tp_size = self.num_gpus
if self.sp_size is None:
self.sp_size = self.num_gpus
if self.tp_size != self.sp_size:
raise ValueError(
f"tp_size ({self.tp_size}) must be equal to sp_size ({self.sp_size})"
)
# Validate VAE spatial parallelism with VAE tiling
if self.vae_sp and not self.vae_tiling:
@@ -414,10 +452,10 @@ class InferenceArgs:
raise ValueError("prompt_path must be a text file")
_inference_args = None
_current_fastvideo_args = None
def prepare_inference_args(argv: List[str]) -> InferenceArgs:
def prepare_fastvideo_args(argv: List[str]) -> FastVideoArgs:
"""
Prepare the inference arguments from the command line arguments.
@@ -429,26 +467,38 @@ def prepare_inference_args(argv: List[str]) -> InferenceArgs:
The inference arguments.
"""
parser = FlexibleArgumentParser()
InferenceArgs.add_cli_args(parser)
FastVideoArgs.add_cli_args(parser)
raw_args = parser.parse_args(argv)
inference_args = InferenceArgs.from_cli_args(raw_args)
inference_args.check_inference_args()
global _inference_args
_inference_args = inference_args
return inference_args
fastvideo_args = FastVideoArgs.from_cli_args(raw_args)
fastvideo_args.check_inference_args()
global _current_fastvideo_args
_current_fastvideo_args = fastvideo_args
return fastvideo_args
def get_inference_args() -> InferenceArgs:
global _inference_args
if _inference_args is None:
raise ValueError("Inference arguments not set")
return _inference_args
@contextmanager
def set_current_fastvideo_args(fastvideo_args: FastVideoArgs):
"""
Temporarily set the current fastvideo config.
Used during model initialization.
We save the current fastvideo config in a global variable,
so that all modules can access it, e.g. custom ops
can access the fastvideo config to determine how to dispatch.
"""
global _current_fastvideo_args
old_fastvideo_args = _current_fastvideo_args
try:
_current_fastvideo_args = fastvideo_args
yield
finally:
_current_fastvideo_args = old_fastvideo_args
class DeprecatedAction(argparse.Action):
def __init__(self, option_strings, dest, nargs=0, **kwargs):
super().__init__(option_strings, dest, nargs=nargs, **kwargs)
def __call__(self, parser, namespace, values, option_string=None):
raise ValueError(self.help)
def get_current_fastvideo_args() -> FastVideoArgs:
if _current_fastvideo_args is None:
# in ci, usually when we test custom ops/modules directly,
# we don't set the fastvideo config. In that case, we set a default
# config.
# TODO(will): may need to handle this for CI.
raise ValueError("Current fastvideo args is not set.")
return _current_fastvideo_args
+2 -2
View File
@@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Optional
import torch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
if TYPE_CHECKING:
@@ -52,7 +52,7 @@ def get_forward_context() -> ForwardContext:
@contextmanager
def set_forward_context(current_timestep,
attn_metadata,
inference_args: InferenceArgs = None):
fastvideo_args: Optional[FastVideoArgs] = None):
"""A context manager that stores the current forward context,
can be attention metadata, etc.
Here we can inject common logic for every model forward pass.
+29 -29
View File
@@ -10,7 +10,7 @@ from typing import Any, Dict
import torch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
build_pipeline)
@@ -28,29 +28,29 @@ class InferenceEngine:
def __init__(
self,
pipeline: ComposedPipelineBase,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
):
"""
Initialize the inference engine.
Args:
pipeline: The pipeline to use for inference.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
default_negative_prompt: The default negative prompt to use.
"""
self.pipeline = pipeline
self.inference_args = inference_args
self.fastvideo_args = fastvideo_args
@classmethod
def create_engine(
cls,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> "InferenceEngine":
"""
Create an inference engine with the specified arguments.
Args:
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
model_loader_cls: The model loader class to use. If None, it will be
determined from the model type.
pipeline_type: The type of pipeline to create. If None, it will be
@@ -71,16 +71,16 @@ class InferenceEngine:
# this way for training we can just do pipeline_cls.from_pretrained(
# checkpoint_path) and have it handle everything.
# TODO(Peiyuan): Then maybe we should only pass in model path and device, not the entire inference args?
pipeline = build_pipeline(inference_args)
pipeline = build_pipeline(fastvideo_args)
logger.info("Pipeline Ready")
# Create the inference engine
return cls(pipeline, inference_args)
return cls(pipeline, fastvideo_args)
def run(
self,
prompt: str,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> Dict[str, Any]:
"""
Run inference with the pipeline.
@@ -96,17 +96,17 @@ class InferenceEngine:
"""
out_dict: Dict[str, Any] = dict()
num_videos_per_prompt = inference_args.num_videos
seed = inference_args.seed
height = inference_args.height
width = inference_args.width
video_length = inference_args.num_frames
negative_prompt = inference_args.neg_prompt
infer_steps = inference_args.num_inference_steps
guidance_scale = inference_args.guidance_scale
flow_shift = inference_args.flow_shift
embedded_guidance_scale = inference_args.embedded_cfg_scale
image_path = inference_args.image_path
num_videos_per_prompt = fastvideo_args.num_videos
seed = fastvideo_args.seed
height = fastvideo_args.height
width = fastvideo_args.width
video_length = fastvideo_args.num_frames
negative_prompt = fastvideo_args.neg_prompt
infer_steps = fastvideo_args.num_inference_steps
guidance_scale = fastvideo_args.guidance_scale
flow_shift = fastvideo_args.flow_shift
embedded_guidance_scale = fastvideo_args.embedded_cfg_scale
image_path = fastvideo_args.image_path
# ========================================================================
# Arguments: target_width, target_height, target_video_length
@@ -161,21 +161,21 @@ class InferenceEngine:
# return
# sp_group = get_sp_group()
# local_rank = sp_group.rank
device = torch.device(inference_args.device_str)
device = torch.device(fastvideo_args.device_str)
batch = ForwardBatch(
image_path=image_path,
prompt=prompt,
negative_prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
height=inference_args.height,
width=inference_args.width,
num_frames=inference_args.num_frames,
num_inference_steps=inference_args.num_inference_steps,
guidance_scale=inference_args.guidance_scale,
height=fastvideo_args.height,
width=fastvideo_args.width,
num_frames=fastvideo_args.num_frames,
num_inference_steps=fastvideo_args.num_inference_steps,
guidance_scale=fastvideo_args.guidance_scale,
# generator=generator,
eta=0.0,
n_tokens=n_tokens,
data_type="video" if inference_args.num_frames > 1 else "image",
data_type="video" if fastvideo_args.num_frames > 1 else "image",
device=device,
extra={}, # Any additional parameters
)
@@ -184,7 +184,7 @@ class InferenceEngine:
print(batch)
print('===============================================')
print('===============================================')
print(inference_args)
print(fastvideo_args)
# ========================================================================
# Pipeline inference
@@ -192,7 +192,7 @@ class InferenceEngine:
start_time = time.time()
samples = self.pipeline.forward(
batch=batch,
inference_args=inference_args,
fastvideo_args=fastvideo_args,
).output
# TODO(will): fix and move to hunyuan stage
# out_dict["seeds"] = batch.seeds
+1 -1
View File
@@ -23,7 +23,7 @@ class SiluAndMul(CustomOp):
return: (num_tokens, d) or (batch_size, seq_len, d)
"""
def __init__(self):
def __init__(self) -> None:
super().__init__()
if current_platform.is_cuda_alike() or current_platform.is_cpu():
self.op = torch.ops._C.silu_and_mul
+3 -1
View File
@@ -113,7 +113,7 @@ class ScaleResidual(nn.Module):
Applies gated residual connection.
"""
def __init__(self):
def __init__(self, prefix: str = ""):
super().__init__()
def forward(self, residual: torch.Tensor, x: torch.Tensor,
@@ -139,6 +139,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
eps: float = 1e-6,
elementwise_affine: bool = False,
dtype: torch.dtype = torch.float32,
prefix: str = "",
):
super().__init__()
if norm_type == "rms":
@@ -189,6 +190,7 @@ class LayerNormScaleShift(nn.Module):
eps: float = 1e-6,
elementwise_affine: bool = False,
dtype: torch.dtype = torch.float32,
prefix: str = "",
):
super().__init__()
if norm_type == "rms":
+4 -3
View File
@@ -778,12 +778,13 @@ class QKVParallelLinear(ColumnParallelLinear):
# no need to narrow
is_sharded_weight = is_sharded_weight
shard_idx = 0
param_data = param_data.narrow(output_dim, shard_offset, shard_size)
if loaded_shard_id == "q":
shard_id = tp_rank
shard_idx = tp_rank
else:
shard_id = tp_rank // self.num_kv_head_replicas
start_idx = shard_id * shard_size
shard_idx = tp_rank // self.num_kv_head_replicas
start_idx = shard_idx * shard_size
if not is_sharded_weight:
loaded_weight = loaded_weight.narrow(output_dim, start_idx,
+1 -1
View File
@@ -12,7 +12,6 @@ from fastvideo.v1.layers.linear import ReplicatedLinear
class MLP(nn.Module):
"""
MLP for DiT blocks, NO gated linear units
TODO: add Tensor Parallel
"""
def __init__(
@@ -23,6 +22,7 @@ class MLP(nn.Module):
bias: bool = True,
act_type: str = "gelu_pytorch_tanh",
dtype: Optional[torch.dtype] = None,
prefix: str = "",
):
super().__init__()
self.fc_in = ReplicatedLinear(
+2 -2
View File
@@ -84,7 +84,7 @@ class RotaryEmbedding(CustomOp):
head_size: int,
rotary_dim: int,
max_position_embeddings: int,
base: int,
base: Union[int, float],
is_neox_style: bool,
dtype: torch.dtype,
) -> None:
@@ -446,7 +446,7 @@ def get_rope(
head_size: int,
rotary_dim: int,
max_position: int,
base: int,
base: Union[int, float],
is_neox_style: bool = True,
rope_scaling: Optional[Dict[str, Any]] = None,
dtype: Optional[torch.dtype] = None,
+5 -2
View File
@@ -32,7 +32,8 @@ class PatchEmbed(nn.Module):
norm_layer=None,
flatten=True,
bias=True,
dtype=None):
dtype=None,
prefix: str = ""):
super().__init__()
# Convert patch_size to 2-tuple
if isinstance(patch_size, (list, tuple)):
@@ -73,6 +74,7 @@ class TimestepEmbedder(nn.Module):
max_period=10000,
dtype=None,
freq_dtype=torch.float32,
prefix: str = "",
):
super().__init__()
self.frequency_embedding_size = frequency_embedding_size
@@ -132,6 +134,7 @@ class ModulateProjection(nn.Module):
factor: int = 2,
act_layer: str = "silu",
dtype: Optional[torch.dtype] = None,
prefix: str = "",
):
super().__init__()
self.factor = factor
@@ -148,7 +151,7 @@ class ModulateProjection(nn.Module):
return x
def unpatchify(x, t, h, w, patch_size, channels):
def unpatchify(x, t, h, w, patch_size, channels) -> torch.Tensor:
"""
Convert patched representation back to image space.
+16 -3
View File
@@ -1,10 +1,12 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import List, Union
from typing import List, Optional, Union
import torch
from torch import nn
from fastvideo.v1.platforms import _Backend
# TODO
class BaseDiT(nn.Module, ABC):
@@ -13,8 +15,9 @@ class BaseDiT(nn.Module, ABC):
_param_names_mapping: dict
hidden_size: int
num_attention_heads: int
_supported_attention_backends: List[_Backend] = []
def __init_subclass__(cls):
def __init_subclass__(cls) -> None:
required_class_attrs = [
"_fsdp_shard_conditions", "_param_names_mapping"
]
@@ -27,20 +30,30 @@ class BaseDiT(nn.Module, ABC):
def __init__(self, *args, **kwargs) -> None:
super().__init__()
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
)
@abstractmethod
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs) -> torch.Tensor:
pass
def __post_init__(self):
def __post_init__(self) -> None:
required_attrs = ["hidden_size", "num_attention_heads"]
for attr in required_attrs:
if not hasattr(self, attr):
raise AttributeError(
f"Subclasses of BaseDiT must define '{attr}' instance variable"
)
@property
def supported_attention_backends(self) -> List[_Backend]:
return self._supported_attention_backends
+127 -88
View File
@@ -19,6 +19,7 @@ from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder,
unpatchify)
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.platforms import _Backend
class HunyuanRMSNorm(nn.Module):
@@ -91,6 +92,8 @@ class MMDoubleStreamBlock(nn.Module):
num_attention_heads: int,
mlp_ratio: float,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[List[_Backend]] = None,
prefix: str = "",
):
super().__init__()
@@ -105,6 +108,7 @@ class MMDoubleStreamBlock(nn.Module):
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.img_mod",
)
# Fused operations for image stream
@@ -123,7 +127,8 @@ class MMDoubleStreamBlock(nn.Module):
self.img_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.img_attn_qkv")
self.img_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.img_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
@@ -131,9 +136,14 @@ class MMDoubleStreamBlock(nn.Module):
self.img_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.img_attn_proj")
self.img_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
self.img_mlp = MLP(hidden_size,
mlp_hidden_dim,
bias=True,
dtype=dtype,
prefix=f"{prefix}.img_mlp")
# Text modulation components
self.txt_mod = ModulateProjection(
@@ -141,6 +151,7 @@ class MMDoubleStreamBlock(nn.Module):
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.txt_mod",
)
# Fused operations for text stream
@@ -173,27 +184,12 @@ class MMDoubleStreamBlock(nn.Module):
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
# Distributed attention
self.attn = DistributedAttention(num_heads=num_attention_heads,
head_size=head_dim,
dropout_rate=0.0,
causal=False)
# QK norm layers for text
self.txt_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype)
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
# Distributed attention
self.attn = DistributedAttention(num_heads=num_attention_heads,
head_size=head_dim,
dropout_rate=0.0,
causal=False)
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn")
def forward(
self,
@@ -303,6 +299,8 @@ class MMSingleStreamBlock(nn.Module):
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[List[_Backend]] = None,
prefix: str = "",
):
super().__init__()
@@ -317,13 +315,15 @@ class MMSingleStreamBlock(nn.Module):
self.linear1 = ReplicatedLinear(hidden_size,
hidden_size * 3 + mlp_hidden_dim,
bias=True,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.linear1")
# Combined projection and MLP output
self.linear2 = ReplicatedLinear(hidden_size + mlp_hidden_dim,
hidden_size,
bias=True,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.linear2")
# QK norm layers
self.q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
@@ -345,13 +345,16 @@ class MMSingleStreamBlock(nn.Module):
self.modulation = ModulateProjection(hidden_size,
factor=3,
act_layer="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.modulation")
# Distributed attention
self.attn = DistributedAttention(num_heads=num_attention_heads,
head_size=head_dim,
dropout_rate=0.0,
causal=False)
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn")
def forward(
self,
@@ -433,6 +436,9 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
lambda n, m: "single" in n and str.isdigit(n.split(".")[-1]),
lambda n, m: "refiner" in n and str.isdigit(n.split(".")[-1]),
]
_supported_attention_backends = [
_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN, _Backend.TORCH_SDPA
]
_param_names_mapping = {
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
@@ -548,24 +554,25 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
}
def __init__(
self,
patch_size: int = 2,
patch_size_t: int = 1,
in_channels: int = 16,
out_channels: int = 16,
num_attention_heads: int = 24,
attention_head_dim: int = 128,
mlp_ratio: float = 4.0,
num_layers: int = 20,
num_single_layers: int = 40,
num_refiner_layers: int = 2,
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56),
guidance_embeds: bool = False,
dtype: Optional[torch.dtype] = None,
text_embed_dim: int = 4096,
pooled_projection_dim: int = 768,
rope_theta: int = 256,
qk_norm: str = "rms_norm", #TODO(PY)
self,
patch_size: int = 2,
patch_size_t: int = 1,
in_channels: int = 16,
out_channels: int = 16,
num_attention_heads: int = 24,
attention_head_dim: int = 128,
mlp_ratio: float = 4.0,
num_layers: int = 20,
num_single_layers: int = 40,
num_refiner_layers: int = 2,
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56),
guidance_embeds: bool = False,
dtype: Optional[torch.dtype] = None,
text_embed_dim: int = 4096,
pooled_projection_dim: int = 768,
rope_theta: int = 256,
qk_norm: str = "rms_norm", #TODO(PY)
prefix="",
):
super().__init__()
hidden_size = attention_head_dim * num_attention_heads
@@ -598,29 +605,35 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
self.img_in = PatchEmbed(self.patch_size,
self.in_channels,
self.hidden_size,
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.img_in")
self.txt_in = SingleTokenRefiner(self.text_states_dim,
hidden_size,
num_attention_heads,
depth=num_refiner_layers,
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.txt_in")
# Time modulation
self.time_in = TimestepEmbedder(self.hidden_size,
act_layer="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.time_in")
# Text modulation
self.vector_in = MLP(self.text_states_dim_2,
self.hidden_size,
self.hidden_size,
act_type="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.vector_in")
# Guidance modulation
self.guidance_in = (TimestepEmbedder(
self.hidden_size, act_layer="silu", dtype=dtype)
self.guidance_in = (TimestepEmbedder(self.hidden_size,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.guidance_in")
if self.guidance_embeds else None)
# Double blocks
@@ -630,7 +643,8 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
) for _ in range(num_layers)
supported_attention_backends=self._supported_attention_backends,
prefix=f"{prefix}.double_blocks.{i}") for i in range(num_layers)
])
# Single blocks
@@ -640,25 +654,29 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
) for _ in range(num_single_layers)
supported_attention_backends=self._supported_attention_backends,
prefix=f"{prefix}.single_blocks.{i+num_layers}")
for i in range(num_single_layers)
])
self.final_layer = FinalLayer(hidden_size,
self.patch_size,
self.out_channels,
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.final_layer")
self.__post_init__()
# TODO: change the input the FORWAD_BACTCH Dict
# TODO: change output to a dict
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
guidance=None,
):
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs):
"""
Forward pass of the HunyuanDiT model.
@@ -760,26 +778,31 @@ class SingleTokenRefiner(nn.Module):
depth=2,
qkv_bias=True,
dtype=None,
prefix: str = "",
) -> None:
super().__init__()
# Input projection
self.input_embedder = ReplicatedLinear(in_channels,
hidden_size,
bias=True,
params_dtype=dtype)
self.input_embedder = ReplicatedLinear(
in_channels,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.input_embedder")
# Timestep embedding
self.t_embedder = TimestepEmbedder(hidden_size,
act_layer="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.t_embedder")
# Context embedding
self.c_embedder = MLP(in_channels,
hidden_size,
hidden_size,
act_type="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.c_embedder")
# Refiner blocks
self.refiner_blocks = nn.ModuleList([
@@ -788,7 +811,8 @@ class SingleTokenRefiner(nn.Module):
num_attention_heads,
qkv_bias=qkv_bias,
dtype=dtype,
) for _ in range(depth)
prefix=f"{prefix}.refiner_blocks.{i}",
) for i in range(depth)
])
def forward(self, x, t):
@@ -822,6 +846,7 @@ class IndividualTokenRefinerBlock(nn.Module):
mlp_ratio=4.0,
qkv_bias=True,
dtype=None,
prefix: str = "",
) -> None:
super().__init__()
self.num_attention_heads = num_attention_heads
@@ -836,12 +861,15 @@ class IndividualTokenRefinerBlock(nn.Module):
self.self_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=qkv_bias,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.self_attn_qkv")
self.self_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=qkv_bias,
params_dtype=dtype)
self.self_attn_proj = ReplicatedLinear(
hidden_size,
hidden_size,
bias=qkv_bias,
params_dtype=dtype,
prefix=f"{prefix}.self_attn_proj")
# MLP
self.norm2 = nn.LayerNorm(hidden_size,
@@ -852,18 +880,25 @@ class IndividualTokenRefinerBlock(nn.Module):
mlp_hidden_dim,
bias=True,
act_type="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.mlp")
# Modulation
self.adaLN_modulation = ModulateProjection(hidden_size,
factor=2,
act_layer="silu",
dtype=dtype)
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
# Scaled dot product attention
self.attn = LocalAttention(
num_heads=num_attention_heads,
head_size=hidden_size // num_attention_heads,
# TODO: remove hardcode; remove STA
supported_attention_backends=[
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA
],
)
def forward(self, x, c):
@@ -902,7 +937,8 @@ class FinalLayer(nn.Module):
hidden_size,
patch_size,
out_channels,
dtype=None) -> None:
dtype=None,
prefix: str = "") -> None:
super().__init__()
# Normalization
@@ -916,13 +952,16 @@ class FinalLayer(nn.Module):
self.linear = ReplicatedLinear(hidden_size,
output_dim,
bias=True,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.linear")
# Modulation
self.adaLN_modulation = ModulateProjection(hidden_size,
factor=2,
act_layer="silu",
dtype=dtype)
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
def forward(self, x, c):
# What the heck HF? Why you change the scale and shift order here???
+40 -32
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any, Dict, List, Optional, Tuple, Union
from typing import List, Optional, Tuple, Union
import torch
import torch.nn as nn
@@ -21,6 +21,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder)
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.platforms import _Backend
class WanImageEmbedding(torch.nn.Module):
@@ -117,7 +118,10 @@ class WanSelfAttention(nn.Module):
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False)
causal=False,
supported_attention_backends=[
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA
])
def forward(self, x: torch.Tensor, context: torch.Tensor,
context_lens: int):
@@ -143,8 +147,8 @@ class WanT2VCrossAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
@@ -158,13 +162,16 @@ class WanT2VCrossAttention(WanSelfAttention):
class WanI2VCrossAttention(WanSelfAttention):
def __init__(self,
dim: int,
num_heads: int,
window_size=(-1, -1),
qk_norm=True,
eps=1e-6) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps)
def __init__(
self,
dim: int,
num_heads: int,
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
supported_attention_backends: Optional[List[str]] = None) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps,
supported_attention_backends)
self.add_k_proj = ReplicatedLinear(dim, dim)
self.add_v_proj = ReplicatedLinear(dim, dim)
@@ -183,11 +190,11 @@ class WanI2VCrossAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
k_img = self.norm_added_k.forward_native(
self.add_k_proj(context_img)[0]).view(b, -1, n, d)
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
b, -1, n, d)
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
img_x = self.attn(q, k_img, v_img)
# compute attention
@@ -212,6 +219,7 @@ class WanTransformerBlock(nn.Module):
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[List[_Backend]] = None,
):
super().__init__()
@@ -221,10 +229,11 @@ class WanTransformerBlock(nn.Module):
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = DistributedAttention(num_heads=num_heads,
head_size=dim // num_heads,
dropout_rate=0.0,
causal=False)
self.attn1 = DistributedAttention(
num_heads=num_heads,
head_size=dim // num_heads,
causal=False,
supported_attention_backends=supported_attention_backends)
self.hidden_dim = dim
self.num_attention_heads = num_heads
dim_head = dim // num_heads
@@ -342,6 +351,7 @@ class WanTransformer3DModel(BaseDiT):
_fsdp_shard_conditions = [
lambda n, m: "blocks" in n and str.isdigit(n.split(".")[-1]),
]
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
_param_names_mapping = {
r"^patch_embedding\.(.*)$":
r"patch_embedding.proj.\1",
@@ -428,7 +438,9 @@ class WanTransformer3DModel(BaseDiT):
self.blocks = nn.ModuleList([
WanTransformerBlock(inner_dim, ffn_dim, num_attention_heads,
qk_norm, cross_attn_norm, eps,
added_kv_proj_dim) for _ in range(num_layers)
added_kv_proj_dim,
self._supported_attention_backends)
for _ in range(num_layers)
])
# 4. Output norm & projection
@@ -446,18 +458,14 @@ class WanTransformer3DModel(BaseDiT):
self.__post_init__()
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
seq_len: Optional[int] = None,
encoder_hidden_states_image: Optional[Union[torch.Tensor,
List[torch.Tensor]]] = None,
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
guidance=None,
) -> torch.Tensor:
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs) -> torch.Tensor:
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
+23
View File
@@ -0,0 +1,23 @@
from typing import List
from torch import nn
from fastvideo.v1.platforms import _Backend
class BaseEncoder(nn.Module):
_supported_attention_backends: List[_Backend] = []
def __init__(self, *args, **kwargs) -> None:
super().__init__()
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
)
def forward(self, *args, **kwargs):
pass
@property
def supported_attention_backends(self) -> List[_Backend]:
return self._supported_attention_backends
+12 -3
View File
@@ -19,11 +19,13 @@ from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
RowParallelLinear)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.encoders.base import BaseEncoder
from fastvideo.v1.models.encoders.vision import (VisionEncoderInfo,
resolve_visual_encoder_outputs)
# TODO: support quantization
# from vllm.model_executor.layers.quantization import QuantizationConfig
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
from fastvideo.v1.platforms import _Backend
logger = init_logger(__name__)
@@ -195,7 +197,9 @@ class CLIPAttention(nn.Module):
self.head_dim,
self.num_heads_per_partition,
softmax_scale=self.scale,
causal=True)
causal=True,
supported_attention_backends=self.config.
supported_attention_backends)
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
return tensor.view(bsz, seq_len, self.num_heads,
@@ -464,7 +468,8 @@ class CLIPTextTransformer(nn.Module):
)
class CLIPTextModel(nn.Module):
class CLIPTextModel(BaseEncoder):
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
def __init__(
self,
@@ -475,6 +480,7 @@ class CLIPTextModel(nn.Module):
super().__init__()
self.config = config
self.config.supported_attention_backends = self._supported_attention_backends
self.text_model = CLIPTextTransformer(config=config,
quant_config=quant_config,
prefix=prefix)
@@ -609,10 +615,11 @@ class CLIPVisionTransformer(nn.Module):
return encoder_outputs
class CLIPVisionModel(nn.Module, SupportsQuant):
class CLIPVisionModel(BaseEncoder, SupportsQuant):
config_class = CLIPVisionConfig
main_input_name = "pixel_values"
packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]}
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
def __init__(
self,
@@ -624,6 +631,8 @@ class CLIPVisionModel(nn.Module, SupportsQuant):
prefix: str = "",
) -> None:
super().__init__()
self.config = config
self.config.supported_attention_backends = self._supported_attention_backends
self.vision_model = CLIPVisionTransformer(
config=config,
quant_config=quant_config,
+13 -8
View File
@@ -39,10 +39,11 @@ from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
QKVParallelLinear, RowParallelLinear)
from fastvideo.v1.layers.rotary_embedding import get_rope
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.v1.models.encoders.base import BaseEncoder
from fastvideo.v1.models.loader.weight_utils import (default_weight_loader,
maybe_remap_kv_scale_name)
# from ..utils import (extract_layer_index)
from fastvideo.v1.platforms import _Backend
class QuantizationConfig:
@@ -159,16 +160,18 @@ class LlamaAttention(nn.Module):
self.head_dim,
rotary_dim=self.rotary_dim,
max_position=max_position_embeddings,
base=rope_theta,
base=int(rope_theta),
rope_scaling=rope_scaling,
is_neox_style=is_neox_style,
)
self.attn = LocalAttention(self.num_heads,
self.head_dim,
self.num_kv_heads,
softmax_scale=self.scaling,
causal=True)
self.attn = LocalAttention(
self.num_heads,
self.head_dim,
self.num_kv_heads,
softmax_scale=self.scaling,
causal=True,
supported_attention_backends=config.supported_attention_backends)
def forward(
self,
@@ -276,7 +279,8 @@ class LlamaDecoderLayer(nn.Module):
return hidden_states, residual
class LlamaModel(nn.Module):
class LlamaModel(BaseEncoder):
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
def __init__(self,
config: LlamaConfig,
@@ -288,6 +292,7 @@ class LlamaModel(nn.Module):
lora_config = None
self.config = config
self.config.supported_attention_backends = self._supported_attention_backends
self.quant_config = quant_config
if lora_config is not None:
max_loras = 1
+2 -3
View File
@@ -52,7 +52,6 @@ def get_hf_config(
trust_remote_code: bool,
revision: Optional[str] = None,
model_override_args: Optional[dict] = None,
inference_args: Optional[dict] = None,
**kwargs,
):
is_gguf = check_gguf_file(model)
@@ -84,13 +83,13 @@ def get_hf_config(
def get_diffusers_config(
model: str,
inference_args: Optional[dict] = None,
fastvideo_args: Optional[dict] = None,
) -> Dict[str, Any]:
"""Gets a configuration for the given diffusers model.
Args:
model: The model name or path.
inference_args: Optional inference arguments to override in the config.
fastvideo_args: Optional inference arguments to override in the config.
Returns:
The loaded configuration.
+32 -34
View File
@@ -13,7 +13,7 @@ from safetensors.torch import load_file as safetensors_load_file
from transformers import AutoImageProcessor, AutoTokenizer, PretrainedConfig
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.hf_transformer_utils import (get_diffusers_config,
get_hf_config)
@@ -36,14 +36,14 @@ class ComponentLoader(ABC):
@abstractmethod
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""
Load the component based on the model path, architecture, and inference args.
Args:
model_path: Path to the component model
architecture: Architecture of the component model
inference_args: Inference arguments
fastvideo_args: Inference arguments
Returns:
The loaded component
@@ -199,21 +199,20 @@ class TextEncoderLoader(ComponentLoader):
yield from self._get_weights_iterator(source)
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the text encoders based on the model path, architecture, and inference args."""
model_config: PretrainedConfig = get_hf_config(
model=model_path,
trust_remote_code=inference_args.trust_remote_code,
revision=inference_args.revision,
trust_remote_code=fastvideo_args.trust_remote_code,
revision=fastvideo_args.revision,
model_override_args=None,
inference_args=inference_args,
)
logger.info("HF Model config: %s", model_config)
target_device = torch.device(inference_args.device_str)
target_device = torch.device(fastvideo_args.device_str)
# TODO(will): add support for other dtypes
return self.load_model(model_path, model_config, target_device,
inference_args.text_encoder_precision)
fastvideo_args.text_encoder_precision)
def load_model(self,
model_path: str,
@@ -250,28 +249,27 @@ class TextEncoderLoader(ComponentLoader):
class ImageEncoderLoader(TextEncoderLoader):
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the text encoders based on the model path, architecture, and inference args."""
model_config: PretrainedConfig = get_hf_config(
model=model_path,
trust_remote_code=inference_args.trust_remote_code,
revision=inference_args.revision,
trust_remote_code=fastvideo_args.trust_remote_code,
revision=fastvideo_args.revision,
model_override_args=None,
inference_args=inference_args,
)
logger.info("HF Model config: %s", model_config)
target_device = torch.device(inference_args.device_str)
target_device = torch.device(fastvideo_args.device_str)
# TODO(will): add support for other dtypes
return self.load_model(model_path, model_config, target_device,
inference_args.image_encoder_precision)
fastvideo_args.image_encoder_precision)
class ImageProcessorLoader(ComponentLoader):
"""Loader for image processor."""
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the image processor based on the model path, architecture, and inference args."""
logger.info("Loading image processor from %s", model_path)
@@ -285,7 +283,7 @@ class TokenizerLoader(ComponentLoader):
"""Loader for tokenizers."""
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the tokenizer based on the model path, architecture, and inference args."""
logger.info("Loading tokenizer from %s", model_path)
@@ -303,7 +301,7 @@ class VAELoader(ComponentLoader):
"""Loader for VAE."""
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the VAE based on the model path, architecture, and inference args."""
# TODO(will): move this to a constants file
config = get_diffusers_config(model=model_path)
@@ -314,7 +312,7 @@ class VAELoader(ComponentLoader):
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(**config).to(inference_args.device)
vae = vae_cls(**config).to(fastvideo_args.device)
# Find all safetensors files
safetensors_list = glob.glob(
@@ -325,7 +323,7 @@ class VAELoader(ComponentLoader):
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
loaded = safetensors_load_file(safetensors_list[0])
vae.load_state_dict(loaded)
dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae = vae.eval().to(dtype)
return vae
@@ -335,7 +333,7 @@ class TransformerLoader(ComponentLoader):
"""Loader for transformer."""
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the transformer based on the model path, architecture, and inference args."""
model_config = get_diffusers_config(model=model_path)
cls_name = model_config.pop("_class_name")
@@ -356,16 +354,16 @@ class TransformerLoader(ComponentLoader):
logger.info("Loading model from %s safetensors files in %s",
len(safetensors_list), model_path)
# initialize_sequence_parallel_group(inference_args.sp_size)
default_dtype = PRECISION_TO_TYPE[inference_args.precision]
# initialize_sequence_parallel_group(fastvideo_args.sp_size)
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
# Load the model using FSDP loader
logger.info("Loading model from %s", cls_name)
model = load_fsdp_model(model_cls=model_cls,
init_params=model_config,
weight_dir_list=safetensors_list,
device=inference_args.device,
cpu_offload=inference_args.use_cpu_offload,
device=fastvideo_args.device,
cpu_offload=fastvideo_args.use_cpu_offload,
default_dtype=default_dtype)
total_params = sum(p.numel() for p in model.parameters())
@@ -382,7 +380,7 @@ class SchedulerLoader(ComponentLoader):
"""Loader for scheduler."""
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the scheduler based on the model path, architecture, and inference args."""
config = get_diffusers_config(model=model_path)
@@ -393,8 +391,8 @@ class SchedulerLoader(ComponentLoader):
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
scheduler = scheduler_cls(**config)
if inference_args.flow_shift is not None:
scheduler.set_shift(inference_args.flow_shift)
if fastvideo_args.flow_shift is not None:
scheduler.set_shift(fastvideo_args.flow_shift)
return scheduler
@@ -407,7 +405,7 @@ class GenericComponentLoader(ComponentLoader):
self.library = library
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load a generic component based on the model path, architecture, and inference args."""
logger.warning("Using generic loader for %s with library %s",
model_path, self.library)
@@ -417,8 +415,8 @@ class GenericComponentLoader(ComponentLoader):
model = AutoModel.from_pretrained(
model_path,
trust_remote_code=inference_args.trust_remote_code,
revision=inference_args.revision,
trust_remote_code=fastvideo_args.trust_remote_code,
revision=fastvideo_args.revision,
)
logger.info("Loaded generic transformers model: %s",
model.__class__.__name__)
@@ -445,7 +443,7 @@ class PipelineComponentLoader:
@staticmethod
def load_module(module_name: str, component_model_path: str,
transformers_or_diffusers: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""
Load a pipeline module.
@@ -454,7 +452,7 @@ class PipelineComponentLoader:
component_model_path: Path to the component model
transformers_or_diffusers: Whether the module is from transformers or diffusers
architecture: Architecture of the component model
inference_args: Inference arguments
fastvideo_args: Inference arguments
Returns:
The loaded module
@@ -471,4 +469,4 @@ class PipelineComponentLoader:
transformers_or_diffusers)
# Load the module
return loader.load(component_model_path, architecture, inference_args)
return loader.load(component_model_path, architecture, fastvideo_args)
+29 -48
View File
@@ -2,7 +2,7 @@
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/parameter.py
from fractions import Fraction
from typing import Any, Callable, Optional, Tuple, Union
from typing import Any, Callable, Tuple, Union
import torch
from torch.nn import Parameter
@@ -58,21 +58,22 @@ class BasevLLMParameter(Parameter):
cond2 = loaded_weight.ndim == 0 and loaded_weight.numel() == 1
return (cond1 and cond2)
def _assert_and_load(self, loaded_weight: torch.Tensor):
def _assert_and_load(self, loaded_weight: torch.Tensor) -> None:
assert (self.data.shape == loaded_weight.shape
or self._is_1d_and_scalar(loaded_weight))
self.data.copy_(loaded_weight)
def load_column_parallel_weight(self, loaded_weight: torch.Tensor):
def load_column_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
self._assert_and_load(loaded_weight)
def load_row_parallel_weight(self, loaded_weight: torch.Tensor):
def load_row_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
self._assert_and_load(loaded_weight)
def load_merged_column_weight(self, loaded_weight: torch.Tensor, **kwargs):
def load_merged_column_weight(self, loaded_weight: torch.Tensor,
**kwargs) -> None:
self._assert_and_load(loaded_weight)
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs):
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs) -> None:
self._assert_and_load(loaded_weight)
@@ -95,7 +96,7 @@ class _ColumnvLLMParameter(BasevLLMParameter):
def output_dim(self):
return self._output_dim
def load_column_parallel_weight(self, loaded_weight: torch.Tensor):
def load_column_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
tp_rank = get_tensor_model_parallel_rank()
shard_size = self.data.shape[self.output_dim]
loaded_weight = loaded_weight.narrow(self.output_dim,
@@ -103,10 +104,13 @@ class _ColumnvLLMParameter(BasevLLMParameter):
assert self.data.shape == loaded_weight.shape
self.data.copy_(loaded_weight)
def load_merged_column_weight(self, loaded_weight: torch.Tensor, **kwargs):
def load_merged_column_weight(self, loaded_weight: torch.Tensor,
**kwargs) -> None:
shard_offset = kwargs.get("shard_offset")
shard_size = kwargs.get("shard_size")
if shard_offset is None or shard_size is None:
raise ValueError("shard_offset and shard_size must be provided")
if isinstance(
self,
(PackedColumnParameter,
@@ -124,13 +128,18 @@ class _ColumnvLLMParameter(BasevLLMParameter):
assert param_data.shape == loaded_weight.shape
param_data.copy_(loaded_weight)
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs):
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs) -> None:
shard_offset = kwargs.get("shard_offset")
shard_size = kwargs.get("shard_size")
shard_id = kwargs.get("shard_id")
num_heads = kwargs.get("num_heads")
assert shard_offset is not None
assert shard_size is not None
assert shard_id is not None
assert num_heads is not None
if isinstance(
self,
(PackedColumnParameter,
@@ -166,7 +175,7 @@ class RowvLLMParameter(BasevLLMParameter):
def input_dim(self):
return self._input_dim
def load_row_parallel_weight(self, loaded_weight: torch.Tensor):
def load_row_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
tp_rank = get_tensor_model_parallel_rank()
shard_size = self.data.shape[self.input_dim]
loaded_weight = loaded_weight.narrow(self.input_dim,
@@ -233,16 +242,16 @@ class PerTensorScaleParameter(BasevLLMParameter):
# For row parallel layers, no sharding needed
# load weight into parameter as is
def load_row_parallel_weight(self, *args, **kwargs):
def load_row_parallel_weight(self, *args, **kwargs) -> None:
super().load_row_parallel_weight(*args, **kwargs)
def load_merged_column_weight(self, *args, **kwargs):
def load_merged_column_weight(self, *args, **kwargs) -> None:
self._load_into_shard_id(*args, **kwargs)
def load_qkv_weight(self, *args, **kwargs):
def load_qkv_weight(self, *args, **kwargs) -> None:
self._load_into_shard_id(*args, **kwargs)
def load_column_parallel_weight(self, *args, **kwargs):
def load_column_parallel_weight(self, *args, **kwargs) -> None:
super().load_row_parallel_weight(*args, **kwargs)
def _load_into_shard_id(self, loaded_weight: torch.Tensor,
@@ -273,14 +282,10 @@ class PackedColumnParameter(_ColumnvLLMParameter):
for more details on the packed properties.
"""
def __init__(self,
packed_factor: Union[int, Fraction],
packed_dim: int,
marlin_tile_size: Optional[int] = None,
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
**kwargs):
self._packed_factor = packed_factor
self._packed_dim = packed_dim
self._marlin_tile_size = marlin_tile_size
super().__init__(**kwargs)
@property
@@ -291,17 +296,12 @@ class PackedColumnParameter(_ColumnvLLMParameter):
def packed_factor(self):
return self._packed_factor
@property
def marlin_tile_size(self):
return self._marlin_tile_size
def adjust_shard_indexes_for_packing(self, shard_size,
shard_offset) -> Tuple[Any, Any]:
return _adjust_shard_indexes_for_packing(
shard_size=shard_size,
shard_offset=shard_offset,
packed_factor=self.packed_factor,
marlin_tile_size=self.marlin_tile_size)
packed_factor=self.packed_factor)
class PackedvLLMParameter(ModelWeightParameter):
@@ -315,14 +315,10 @@ class PackedvLLMParameter(ModelWeightParameter):
by accounting for packing and optionally, marlin tile size.
"""
def __init__(self,
packed_factor: Union[int, Fraction],
packed_dim: int,
marlin_tile_size: Optional[int] = None,
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
**kwargs):
self._packed_factor = packed_factor
self._packed_dim = packed_dim
self._marlin_tile_size = marlin_tile_size
super().__init__(**kwargs)
@property
@@ -333,16 +329,11 @@ class PackedvLLMParameter(ModelWeightParameter):
def packed_factor(self):
return self._packed_factor
@property
def marlin_tile_size(self):
return self._marlin_tile_size
def adjust_shard_indexes_for_packing(self, shard_size, shard_offset):
return _adjust_shard_indexes_for_packing(
shard_size=shard_size,
shard_offset=shard_offset,
packed_factor=self.packed_factor,
marlin_tile_size=self.marlin_tile_size)
packed_factor=self.packed_factor)
class BlockQuantScaleParameter(_ColumnvLLMParameter, RowvLLMParameter):
@@ -412,18 +403,8 @@ def permute_param_layout_(param: BasevLLMParameter, input_dim: int,
return param
def _adjust_shard_indexes_for_marlin(shard_size, shard_offset,
marlin_tile_size) -> Tuple[Any, Any]:
return shard_size * marlin_tile_size, shard_offset * marlin_tile_size
def _adjust_shard_indexes_for_packing(shard_size, shard_offset, packed_factor,
marlin_tile_size) -> Tuple[Any, Any]:
def _adjust_shard_indexes_for_packing(shard_size, shard_offset,
packed_factor) -> Tuple[Any, Any]:
shard_size = shard_size // packed_factor
shard_offset = shard_offset // packed_factor
if marlin_tile_size is not None:
return _adjust_shard_indexes_for_marlin(
shard_size=shard_size,
shard_offset=shard_offset,
marlin_tile_size=marlin_tile_size)
return shard_size, shard_offset
+4 -5
View File
@@ -8,7 +8,7 @@ from diffusers.utils import BaseOutput
class BaseScheduler(ABC):
timesteps: torch.tensor
timesteps: torch.Tensor
order: int
def __init__(self, *args, **kwargs) -> None:
@@ -38,10 +38,9 @@ class BaseScheduler(ABC):
@abstractmethod
def step(
self,
model_output: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
model_output: torch.Tensor,
timestep: Union[int, torch.Tensor],
sample: torch.Tensor,
return_dict: bool = True,
**kwargs,
) -> Union[BaseOutput, Tuple]:
pass
@@ -200,6 +200,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
return_dict: bool = True,
**kwargs,
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
-202
View File
@@ -1,202 +0,0 @@
from dataclasses import dataclass
from typing import Any, Optional
import torch
import torch.nn as nn
from transformers.utils import ModelOutput
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def use_default(value, default) -> Any:
return value if value is not None else default
@dataclass
class TextEncoderModelOutput(ModelOutput):
"""
Base class for model's outputs that also contains a pooling of the last hidden states.
Args:
hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
Sequence of hidden-states at the output of the last layer of the model.
attention_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
Mask to avoid performing attention on padding token indices. Mask values selected in ``[0, 1]``:
hidden_states_list (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed):
Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
text_outputs (`list`, *optional*, returned when `return_texts=True` is passed):
List of decoded texts.
"""
hidden_state: torch.FloatTensor = None
attention_mask: Optional[torch.LongTensor] = None
text_outputs: Optional[list] = None
class TextEncoder(nn.Module):
def __init__(
self,
text_encoder,
tokenizer,
max_length: int,
text_encoder_precision: Optional[str] = None,
text_encoder_path: Optional[str] = None,
output_key: Optional[str] = None,
use_attention_mask: bool = True,
prompt_template: Optional[dict] = None,
prompt_template_video: Optional[dict] = None,
hidden_state_skip_layer: Optional[int] = None,
apply_final_norm: bool = False,
device=None,
):
super().__init__()
# TODO(will): check if there's a cleaner way to do this
self.text_encoder_type = text_encoder.config.architectures[0]
self.max_length = max_length
self.precision = text_encoder_precision
self.model_path = text_encoder_path
self.use_attention_mask = use_attention_mask
if prompt_template_video is not None:
assert (use_attention_mask is True
), "Attention mask is True required when training videos."
self.prompt_template = prompt_template
self.prompt_template_video = prompt_template_video
self.hidden_state_skip_layer = hidden_state_skip_layer
self.apply_final_norm = apply_final_norm
if "T5" in self.text_encoder_type:
self.output_key = output_key or "last_hidden_state"
elif "CLIPTextModel" in self.text_encoder_type:
self.output_key = output_key or "pooler_output"
elif "LlamaModel" in self.text_encoder_type or "glm" in self.text_encoder_type:
self.output_key = output_key or "last_hidden_state"
else:
raise ValueError(
f"Unsupported text encoder type: {self.text_encoder_type}")
self.model = text_encoder
# self.dtype = self.model.dtype
self.device = device
self.tokenizer = tokenizer
def __repr__(self):
return f"{self.text_encoder_type} ({self.precision} - {self.model_path})"
@staticmethod
def apply_text_to_template(text, template, prevent_empty_text=True) -> str:
"""
Apply text to template.
Args:
text (str): Input text.
template (str or list): Template string or list of chat conversation.
prevent_empty_text (bool): If True, we will prevent the user text from being empty
by adding a space. Defaults to True.
"""
if isinstance(template, str):
# Will send string to tokenizer. Used for llm
return template.format(text)
else:
raise TypeError(f"Unsupported template type: {type(template)}")
def text2tokens(self, text) -> dict:
"""
Tokenize the input text.
Args:
text (str or list): Input text.
"""
if self.prompt_template_video is not None:
prompt_template = self.prompt_template_video["template"]
text = self.apply_text_to_template(text, prompt_template)
kwargs = dict(
truncation=True,
max_length=self.max_length,
return_tensors="pt",
)
batch_encoding: dict = self.tokenizer(
text,
return_length=False,
return_overflowing_tokens=False,
return_attention_mask=True,
**kwargs,
)
return batch_encoding
def encode(
self,
batch_encoding,
use_attention_mask=None,
hidden_state_skip_layer=None,
device=None,
) -> TextEncoderModelOutput:
"""
Args:
batch_encoding (dict): Batch encoding from tokenizer.
use_attention_mask (bool): Whether to use attention mask. If None, use self.use_attention_mask.
Defaults to None.
output_hidden_states (bool): Whether to output hidden states. If False, return the value of
self.output_key. If True, return the entire output. If set self.hidden_state_skip_layer,
output_hidden_states will be set True. Defaults to False.
hidden_state_skip_layer (int): Number of hidden states to hidden_state_skip_layer. 0 means the last layer.
If None, self.output_key will be used. Defaults to None.
return_texts (bool): Whether to return the decoded texts. Defaults to False.
"""
device = self.model.device if device is None else device
use_attention_mask = use_default(use_attention_mask,
self.use_attention_mask)
hidden_state_skip_layer = use_default(hidden_state_skip_layer,
self.hidden_state_skip_layer)
# note: clip will need attention mask
# TODO(will): unify interface with dit
# TODO (peiyuan): why clip need attention mask?
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs = self.model(
input_ids=batch_encoding["input_ids"].to(device),
output_hidden_states=hidden_state_skip_layer is not None,
)
if hidden_state_skip_layer is not None:
last_hidden_state = outputs.hidden_states[-(
hidden_state_skip_layer + 1)]
# Real last hidden state already has layer norm applied. So here we only apply it
# for intermediate layers.
if hidden_state_skip_layer > 0 and self.apply_final_norm:
last_hidden_state = self.model.final_layer_norm(
last_hidden_state)
else:
last_hidden_state = outputs[self.output_key]
# Remove hidden states of instruction tokens, only keep prompt tokens.
if self.prompt_template_video is not None:
crop_start = self.prompt_template_video.get("crop_start", -1)
last_hidden_state = last_hidden_state[:, crop_start:]
return TextEncoderModelOutput(last_hidden_state)
def forward(
self,
text,
use_attention_mask=None,
output_hidden_states=False,
hidden_state_skip_layer=None,
return_texts=False,
):
batch_encoding = self.text2tokens(text)
return self.encode(
batch_encoding,
use_attention_mask=use_attention_mask,
hidden_state_skip_layer=hidden_state_skip_layer,
)
+294 -33
View File
@@ -19,14 +19,34 @@ from typing import Optional, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.checkpoint
from contextlib import contextmanager
import contextvars
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.models.utils import auto_attributes
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.models.vaes.common import ParallelTiledVAE, DiagonalGaussianDistribution
CACHE_T = 2
is_first_frame = contextvars.ContextVar("is_first_frame", default=False)
feat_cache = contextvars.ContextVar("feat_cache", default=None)
feat_idx = contextvars.ContextVar("feat_idx", default=0)
@contextmanager
def forward_context(first_frame_arg=False,
feat_cache_arg=None,
feat_idx_arg=None):
is_first_frame_token = is_first_frame.set(first_frame_arg)
feat_cache_token = feat_cache.set(feat_cache_arg)
feat_idx_token = feat_idx.set(feat_idx_arg)
try:
yield
finally:
is_first_frame.reset(is_first_frame_token)
feat_cache.reset(feat_cache_token)
feat_idx.reset(feat_idx_token)
class WanCausalConv3d(nn.Conv3d):
r"""
@@ -60,12 +80,17 @@ class WanCausalConv3d(nn.Conv3d):
)
self.padding: Tuple[int, int, int]
# Set up causal padding
self._padding = (self.padding[2], self.padding[2], self.padding[1],
self.padding[1], 2 * self.padding[0], 0)
self._padding: Tuple[int, ...] = (self.padding[2], self.padding[2],
self.padding[1], self.padding[1],
2 * self.padding[0], 0)
self.padding = (0, 0, 0)
def forward(self, x):
def forward(self, x, cache_x=None):
padding = list(self._padding)
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
x = torch.cat([cache_x, x], dim=2)
padding[4] -= cache_x.shape[2]
x = F.pad(x, padding)
return super().forward(x)
@@ -157,28 +182,82 @@ class WanResample(nn.Module):
self.time_conv = WanCausalConv3d(dim,
dim, (3, 1, 1),
stride=(2, 1, 1),
padding=(1, 0, 0))
padding=(0, 0, 0))
else:
self.resample = nn.Identity()
def forward(self, x, first_frame=False):
def forward(self, x):
b, c, t, h, w = x.size()
first_frame = is_first_frame.get()
if first_frame:
assert t == 1
if self.mode == "upsample3d" and not first_frame and hasattr(
self, "time_conv"):
x = self.time_conv(x)
x = x.reshape(b, 2, c, t, h, w)
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3)
x = x.reshape(b, c, t * 2, h, w)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if self.mode == "upsample3d":
if _feat_cache is not None:
idx = _feat_idx
if _feat_cache[idx] is None:
_feat_cache[idx] = "Rep"
_feat_idx += 1
else:
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[
idx] is not None and _feat_cache[idx] != "Rep":
# cache last frame of last two chunk
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
if cache_x.shape[2] < 2 and _feat_cache[
idx] is not None and _feat_cache[idx] == "Rep":
cache_x = torch.cat([
torch.zeros_like(cache_x).to(cache_x.device),
cache_x
],
dim=2)
if _feat_cache[idx] == "Rep":
x = self.time_conv(x)
else:
x = self.time_conv(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
x = x.reshape(b, 2, c, t, h, w)
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]),
3)
x = x.reshape(b, c, t * 2, h, w)
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
elif not first_frame and hasattr(self, "time_conv"):
x = self.time_conv(x)
x = x.reshape(b, 2, c, t, h, w)
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3)
x = x.reshape(b, c, t * 2, h, w)
t = x.shape[2]
x = x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)
x = self.resample(x)
x = x.view(b, t, x.size(1), x.size(2), x.size(3)).permute(0, 2, 1, 3, 4)
if self.mode == "downsample3d" and not first_frame and hasattr(
self, "time_conv"):
x = self.time_conv(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if self.mode == "downsample3d":
if _feat_cache is not None:
idx = _feat_idx
if _feat_cache[idx] is None:
_feat_cache[idx] = x.clone()
_feat_idx += 1
else:
cache_x = x[:, :, -1:, :, :].clone()
x = self.time_conv(
torch.cat([_feat_cache[idx][:, :, -1:, :, :], x], 2))
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
elif not first_frame and hasattr(self, "time_conv"):
x = self.time_conv(x)
return x
@@ -222,7 +301,25 @@ class WanResidualBlock(nn.Module):
x = self.norm1(x)
x = self.nonlinearity(x)
x = self.conv1(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv1(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv1(x)
# Second normalization and activation
x = self.norm2(x)
@@ -231,7 +328,25 @@ class WanResidualBlock(nn.Module):
# Dropout
x = self.dropout(x)
x = self.conv2(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv2(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv2(x)
# Add residual connection
return x + h
@@ -400,15 +515,30 @@ class WanEncoder3d(nn.Module):
self.gradient_checkpointing = False
def forward(self, x, first_frame=False):
x = self.conv_in(x)
def forward(self, x):
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv_in(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv_in(x)
## downsamples
for layer in self.down_blocks:
if isinstance(layer, WanResample):
x = layer(x, first_frame=first_frame)
else:
x = layer(x)
x = layer(x)
## middle
x = self.mid_block(x)
@@ -416,7 +546,26 @@ class WanEncoder3d(nn.Module):
## head
x = self.norm_out(x)
x = self.nonlinearity(x)
x = self.conv_out(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv_out(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv_out(x)
return x
@@ -465,7 +614,7 @@ class WanUpBlock(nn.Module):
self.gradient_checkpointing = False
def forward(self, x, first_frame=False):
def forward(self, x):
"""
Forward pass through the upsampling block.
@@ -481,7 +630,7 @@ class WanUpBlock(nn.Module):
x = resnet(x)
if self.upsamplers is not None:
x = self.upsamplers[0](x, first_frame=first_frame)
x = self.upsamplers[0](x)
return x
@@ -569,21 +718,57 @@ class WanDecoder3d(nn.Module):
self.gradient_checkpointing = False
def forward(self, x, first_frame=False):
def forward(self, x):
## conv1
x = self.conv_in(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv_in(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv_in(x)
## middle
x = self.mid_block(x)
## upsamples
for up_block in self.up_blocks:
x = up_block(x, first_frame=first_frame)
x = up_block(x)
## head
x = self.norm_out(x)
x = self.nonlinearity(x)
x = self.conv_out(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv_out(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv_out(x)
return x
@@ -681,10 +866,63 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
self.tile_sample_stride_height = 192
self.tile_sample_stride_width = 192
self.tile_sample_stride_num_frames = 12
# Whether to use the feature cache algorithm used by diffusers and Wan2.1
self.use_feature_cache = True # default to True for best performance
ParallelTiledVAE.__init__(self)
def clear_cache(self) -> None:
def _count_conv3d(model) -> int:
count = 0
for m in model.modules():
if isinstance(m, WanCausalConv3d):
count += 1
return count
self._conv_num = _count_conv3d(self.decoder)
self._conv_idx = 0
self._feat_map = [None] * self._conv_num
# cache encode
self._enc_conv_num = _count_conv3d(self.encoder)
self._enc_conv_idx = 0
self._enc_feat_map = [None] * self._enc_conv_num
def encode(self, x: torch.Tensor) -> torch.Tensor:
if self.use_feature_cache:
self.clear_cache()
with forward_context(feat_cache_arg=self._enc_feat_map,
feat_idx_arg=self._enc_conv_idx):
t = x.shape[2]
iter_ = 1 + (t - 1) // 4
for i in range(iter_):
feat_idx.set(0)
if i == 0:
out = self.encoder(x[:, :, :1, :, :])
else:
out_ = self.encoder(x[:, :,
1 + 4 * (i - 1):1 + 4 * i, :, :])
out = torch.cat([out, out_], 2)
enc = self.quant_conv(out)
mu, logvar = enc[:, :self.z_dim, :, :, :], enc[:,
self.z_dim:, :, :, :]
enc = torch.cat([mu, logvar], dim=1)
enc = DiagonalGaussianDistribution(enc)
self.clear_cache()
else:
for block in self.encoder.down_blocks:
if isinstance(block,
WanResample) and block.mode == "downsample3d":
_padding = list(block.time_conv._padding)
_padding[4] = 2
block.time_conv._padding = tuple(_padding)
enc = ParallelTiledVAE.encode(self, x)
return enc
def _encode(self, x: torch.Tensor, first_frame=False) -> torch.Tensor:
out = self.encoder(x, first_frame=first_frame)
with forward_context(first_frame_arg=first_frame):
out = self.encoder(x)
enc = self.quant_conv(out)
mu, logvar = enc[:, :self.z_dim, :, :, :], enc[:, self.z_dim:, :, :, :]
enc = torch.cat([mu, logvar], dim=1)
@@ -708,9 +946,32 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
enc = torch.cat([first_frame, enc], dim=2)
return enc
def decode(self, z: torch.Tensor) -> torch.Tensor:
if self.use_feature_cache:
self.clear_cache()
iter_ = z.shape[2]
x = self.post_quant_conv(z)
with forward_context(feat_cache_arg=self._feat_map,
feat_idx_arg=self._conv_idx):
for i in range(iter_):
feat_idx.set(0)
if i == 0:
out = self.decoder(x[:, :, i:i + 1, :, :])
else:
out_ = self.decoder(x[:, :, i:i + 1, :, :])
out = torch.cat([out, out_], 2)
out = torch.clamp(out, min=-1.0, max=1.0)
self.clear_cache()
else:
out = ParallelTiledVAE.decode(self, z)
return out
def _decode(self, z: torch.Tensor, first_frame=False) -> torch.Tensor:
x = self.post_quant_conv(z)
out = self.decoder(x, first_frame=first_frame)
with forward_context(first_frame_arg=first_frame):
out = self.decoder(x)
out = torch.clamp(out, min=-1.0, max=1.0)
+5 -5
View File
@@ -40,7 +40,7 @@ from fastvideo.v1.pipelines.stages import (
ConditioningStage,
# Import other required stages
)
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
class YourCustomPipeline(ComposedPipelineBase):
@@ -53,7 +53,7 @@ class YourCustomPipeline(ComposedPipelineBase):
# Add other required modules
]
def create_pipeline_stages(self, inference_args: InferenceArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
# Add and configure pipeline stages
self.add_stage(
stage_name="input_validation_stage",
@@ -61,14 +61,14 @@ class YourCustomPipeline(ComposedPipelineBase):
)
# Add more stages as needed
def initialize_pipeline(self, inference_args: InferenceArgs):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Initialize pipeline-specific components
pass
@torch.no_grad()
def forward(self, batch: ForwardBatch, inference_args: InferenceArgs) -> ForwardBatch:
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
# Implement your pipeline's forward pass
batch = self.input_validation_stage(batch, inference_args)
batch = self.input_validation_stage(batch, fastvideo_args)
# Add more stage executions
return batch
+5 -5
View File
@@ -5,7 +5,7 @@ Diffusion pipelines for fastvideo.v1.
This package contains diffusion pipelines for generating videos and images.
"""
from fastvideo.v1.inference_args import InferenceArgs
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.pipeline_batch_info import ForwardBatch
@@ -16,7 +16,7 @@ from fastvideo.v1.utils import (maybe_download_model,
logger = init_logger(__name__)
def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
def build_pipeline(fastvideo_args: FastVideoArgs) -> ComposedPipelineBase:
"""
Only works with valid hf diffusers configs. (model_index.json)
We want to build a pipeline based on the inference args mode_path:
@@ -25,9 +25,9 @@ def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
3. based on the config, determine the pipeline class
"""
# Get pipeline type
model_path = inference_args.model_path
model_path = fastvideo_args.model_path
model_path = maybe_download_model(model_path)
# inference_args.downloaded_model_path = model_path
# fastvideo_args.downloaded_model_path = model_path
logger.info("Model path: %s", model_path)
config = verify_model_config_and_directory(model_path)
@@ -41,7 +41,7 @@ def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
pipeline_architecture)
# instantiate the pipeline
pipeline = pipeline_cls(model_path, inference_args, config)
pipeline = pipeline_cls(model_path, fastvideo_args, config)
logger.info("Pipeline instantiated")
# pipeline is now initialized and ready to use
@@ -12,7 +12,7 @@ from typing import Any, Dict, List, Optional, cast
import torch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -38,7 +38,7 @@ class ComposedPipelineBase(ABC):
# TODO(will): args should support both inference args and training args
def __init__(self,
model_path: str,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
config: Optional[Dict[str, Any]] = None):
"""
Initialize the pipeline. After __init__, the pipeline should be ready to
@@ -61,12 +61,12 @@ class ComposedPipelineBase(ABC):
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
self.modules = self.load_modules(inference_args)
self.modules = self.load_modules(fastvideo_args)
self.initialize_pipeline(inference_args)
self.initialize_pipeline(fastvideo_args)
logger.info("Creating pipeline stages...")
self.create_pipeline_stages(inference_args)
self.create_pipeline_stages(fastvideo_args)
def get_module(self, module_name: str) -> Any:
return self.modules[module_name]
@@ -77,7 +77,7 @@ class ComposedPipelineBase(ABC):
def _load_config(self, model_path: str) -> Dict[str, Any]:
model_path = maybe_download_model(self.model_path)
self.model_path = model_path
# inference_args.downloaded_model_path = model_path
# fastvideo_args.downloaded_model_path = model_path
logger.info("Model path: %s", model_path)
config = verify_model_config_and_directory(model_path)
return cast(Dict[str, Any], config)
@@ -108,20 +108,20 @@ class ComposedPipelineBase(ABC):
return self._stages
@abstractmethod
def create_pipeline_stages(self, inference_args: InferenceArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""
Create the pipeline stages.
"""
raise NotImplementedError
@abstractmethod
def initialize_pipeline(self, inference_args: InferenceArgs):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
raise NotImplementedError
def load_modules(self, inference_args: InferenceArgs) -> Dict[str, Any]:
def load_modules(self, fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
"""
Load the modules from the config.
"""
@@ -156,7 +156,7 @@ class ComposedPipelineBase(ABC):
component_model_path=component_model_path,
transformers_or_diffusers=transformers_or_diffusers,
architecture=architecture,
inference_args=inference_args,
fastvideo_args=fastvideo_args,
)
logger.info("Loaded module %s from %s", module_name,
component_model_path)
@@ -185,14 +185,14 @@ class ComposedPipelineBase(ABC):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Generate a video or image using the pipeline.
Args:
batch: The batch to generate from.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
ForwardBatch: The batch with the generated video or image.
"""
@@ -201,7 +201,7 @@ class ComposedPipelineBase(ABC):
self._stage_name_mapping.keys())
logger.info("Batch: %s", batch)
for stage in self.stages:
batch = stage(batch, inference_args)
batch = stage(batch, fastvideo_args)
# Return the output
return batch
@@ -8,7 +8,7 @@ using the modular pipeline architecture.
from diffusers.image_processor import VaeImageProcessor
from fastvideo.v1.inference_args import InferenceArgs
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.stages import (CLIPTextEncodingStage,
@@ -30,7 +30,7 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
"transformer", "scheduler"
]
def create_pipeline_stages(self, inference_args: InferenceArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
@@ -67,20 +67,20 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def initialize_pipeline(self, inference_args: InferenceArgs):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
vae_scale_factor = 2**(len(self.get_module("vae").block_out_channels) -
1)
inference_args.vae_scale_factor = vae_scale_factor
fastvideo_args.vae_scale_factor = vae_scale_factor
self.image_processor = VaeImageProcessor(
vae_scale_factor=vae_scale_factor)
self.add_module("image_processor", self.image_processor)
num_channels_latents = self.get_module("transformer").in_channels
inference_args.num_channels_latents = num_channels_latents
fastvideo_args.num_channels_latents = num_channels_latents
EntryClass = HunyuanVideoPipeline
@@ -22,7 +22,7 @@ class ForwardBatch:
execution, allowing methods to update specific components without needing
to manage numerous individual parameters.
"""
# TODO(will): double check that args are separate from inference_args
# TODO(will): double check that args are separate from fastvideo_args
# properly. Also maybe think about providing an abstraction for pipeline
# specific arguments.
data_type: str
@@ -40,8 +40,6 @@ class ForwardBatch:
# Primary encoder embeddings
prompt_embeds: List[torch.Tensor] = field(default_factory=list)
negative_prompt_embeds: Optional[List[torch.Tensor]] = None
attention_mask: List[torch.Tensor] = field(default_factory=list)
negative_attention_mask: List[torch.Tensor] = field(default_factory=list)
# Additional text-related parameters
max_sequence_length: Optional[int] = None
+8 -8
View File
@@ -12,7 +12,7 @@ from abc import ABC, abstractmethod
import torch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -45,7 +45,7 @@ class PipelineStage(ABC):
def __call__(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Execute the stage's processing on the batch with optional logging.
@@ -53,7 +53,7 @@ class PipelineStage(ABC):
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The updated batch information after this stage's processing.
@@ -65,7 +65,7 @@ class PipelineStage(ABC):
try:
# Call the actual implementation
result = self._call_implementation(batch, inference_args)
result = self._call_implementation(batch, fastvideo_args)
execution_time = time.time() - start_time
self._logger.info("[%s] Execution completed in %s ms",
@@ -85,13 +85,13 @@ class PipelineStage(ABC):
else:
# Just call the implementation directly if logging is disabled
# TODO(will): Also handle backward
return self.forward(batch, inference_args)
return self.forward(batch, fastvideo_args)
@abstractmethod
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Forward pass of the stage's processing.
@@ -101,7 +101,7 @@ class PipelineStage(ABC):
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The updated batch information after this stage's processing.
@@ -111,6 +111,6 @@ class PipelineStage(ABC):
def backward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
raise NotImplementedError
@@ -8,7 +8,7 @@ This module contains implementations of image encoding stages for diffusion pipe
import torch
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.inference_args import InferenceArgs
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
@@ -40,19 +40,19 @@ class CLIPImageEncodingStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode the prompt into image encoder hidden states.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded prompt embeddings.
"""
if inference_args.use_cpu_offload:
if fastvideo_args.use_cpu_offload:
self.image_encoder = self.image_encoder.to(batch.device)
image = load_image(batch.image_path)
@@ -64,7 +64,7 @@ class CLIPImageEncodingStage(PipelineStage):
batch.image_embeds.append(image_embeds)
if inference_args.use_cpu_offload:
if fastvideo_args.use_cpu_offload:
self.image_encoder.to('cpu')
torch.cuda.empty_cache()
@@ -8,7 +8,7 @@ This module contains implementations of prompt encoding stages for diffusion pip
import torch
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.inference_args import InferenceArgs
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
@@ -39,19 +39,19 @@ class CLIPTextEncodingStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode the prompt into text encoder hidden states.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded prompt embeddings.
"""
if inference_args.use_cpu_offload:
if fastvideo_args.use_cpu_offload:
self.text_encoder = self.text_encoder.to(batch.device)
text_inputs = self.tokenizer(
@@ -82,9 +82,10 @@ class CLIPTextEncodingStage(PipelineStage):
batch.device), )
negative_prompt_embeds = negative_outputs["pooler_output"]
assert batch.negative_prompt_embeds is not None
batch.negative_prompt_embeds.append(negative_prompt_embeds)
if inference_args.use_cpu_offload:
if fastvideo_args.use_cpu_offload:
self.text_encoder.to('cpu')
torch.cuda.empty_cache()
@@ -5,7 +5,7 @@ Conditioning stage for diffusion pipelines.
import torch
from fastvideo.v1.inference_args import InferenceArgs
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
@@ -24,14 +24,14 @@ class ConditioningStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Apply conditioning to the diffusion process.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with applied conditioning.
+8 -8
View File
@@ -5,7 +5,7 @@ Decoding stage for diffusion pipelines.
import torch
from fastvideo.v1.inference_args import InferenceArgs
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
@@ -28,14 +28,14 @@ class DecodingStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Decode latent representations into pixel space.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with decoded outputs.
@@ -46,13 +46,13 @@ class DecodingStage(PipelineStage):
raise ValueError("Latents must be provided")
# Skip decoding if output type is latent
if inference_args.output_type == "latent":
if fastvideo_args.output_type == "latent":
image = latents
else:
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae_autocast_enabled = (vae_dtype != torch.float32
) and not inference_args.disable_autocast
) and not fastvideo_args.disable_autocast
if isinstance(self.vae.scaling_factor, torch.Tensor):
latents = latents / self.vae.scaling_factor.to(
@@ -73,9 +73,9 @@ class DecodingStage(PipelineStage):
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if inference_args.vae_tiling:
if fastvideo_args.vae_tiling:
self.vae.enable_tiling()
# if inference_args.vae_sp:
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
if not vae_autocast_enabled:
latents = latents.to(vae_dtype)
+32 -26
View File
@@ -3,6 +3,7 @@
Denoising stage for diffusion pipelines.
"""
import importlib.util
import inspect
from typing import Any, Dict, Iterable, Optional
@@ -16,12 +17,21 @@ from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
from fastvideo.v1.distributed.communication_op import (
sequence_model_parallel_all_gather)
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.inference_args import InferenceArgs
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.platforms import _Backend
from fastvideo.v1.utils import PRECISION_TO_TYPE
st_attn_available = False
spec = importlib.util.find_spec("st_attn")
if spec is not None:
st_attn_available = True
from fastvideo.v1.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
logger = init_logger(__name__)
@@ -41,20 +51,20 @@ class DenoisingStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Run the denoising loop.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with denoised latents.
"""
# If use cpu offload, need to load the model back into gpu again
if inference_args.use_cpu_offload:
if fastvideo_args.use_cpu_offload:
self.transformer = self.transformer.to(batch.device)
# Prepare extra step kwargs for scheduler
extra_step_kwargs = self.prepare_extra_func_kwargs(
@@ -66,9 +76,9 @@ class DenoisingStage(PipelineStage):
)
# Setup precision and autocast settings
target_dtype = PRECISION_TO_TYPE[inference_args.precision]
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
autocast_enabled = (target_dtype != torch.float32
) and not inference_args.disable_autocast
) and not fastvideo_args.disable_autocast
# Handle sequence parallelism if enabled
world_size, rank = get_sequence_model_parallel_world_size(
@@ -128,6 +138,7 @@ class DenoisingStage(PipelineStage):
assert torch.isnan(prompt_embeds[0]).sum() == 0
if batch.do_classifier_free_guidance:
neg_prompt_embeds = batch.negative_prompt_embeds
assert neg_prompt_embeds is not None
assert torch.isnan(neg_prompt_embeds[0]).sum() == 0
# Run denoising loop
@@ -150,11 +161,11 @@ class DenoisingStage(PipelineStage):
# Prepare inputs for transformer
t_expand = t.repeat(latent_model_input.shape[0])
guidance_expand = (torch.tensor(
[inference_args.embedded_cfg_scale] *
[fastvideo_args.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=batch.device,
).to(target_dtype) * 1000.0 if inference_args.embedded_cfg_scale
).to(target_dtype) * 1000.0 if fastvideo_args.embedded_cfg_scale
is not None else None)
# Predict noise residual
@@ -167,41 +178,36 @@ class DenoisingStage(PipelineStage):
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
distributed=True,
supported_attention_backends=[
_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN,
_Backend.TORCH_SDPA
] # hack
)
# TODO(will): clean this up...
try:
from fastvideo.v1.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
except ImportError:
SlidingTileAttentionBackend = None
if SlidingTileAttentionBackend is not None and isinstance(
self.attn_backend, SlidingTileAttentionBackend):
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls(
)
# TODO(will-refactor): should this be in a new stage?
# TODO(will): clean this up
attn_metadata = self.attn_metadata_builder.build(
current_timestep=i,
forward_batch=batch,
inference_args=inference_args,
fastvideo_args=fastvideo_args,
)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
# TODO(will): finalize the interface. vLLM uses this to
# support torch dynamo compilation. They pass in
# attn_metadata, vllm_config, and num_tokens. We can pass in
# inference_args or training_args, and attn_metadata.
# fastvideo_args or training_args, and attn_metadata.
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
# inference_args=inference_args
# fastvideo_args=fastvideo_args
):
# Run transformer
noise_pred = self.transformer(
@@ -217,7 +223,7 @@ class DenoisingStage(PipelineStage):
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
# inference_args=inference_args
# fastvideo_args=fastvideo_args
):
# Run transformer
noise_pred_uncond = self.transformer(
@@ -261,7 +267,7 @@ class DenoisingStage(PipelineStage):
# Update batch with final latents
batch.latents = latents
if inference_args.use_cpu_offload:
if fastvideo_args.use_cpu_offload:
self.transformer.to('cpu')
torch.cuda.empty_cache()
+12 -10
View File
@@ -7,7 +7,7 @@ from typing import Optional
import PIL.Image
import torch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vision_utils import (get_default_height_width,
load_image, normalize,
@@ -33,14 +33,14 @@ class EncodingStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode pixel representations into latent space.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded outputs.
@@ -62,7 +62,7 @@ class EncodingStage(PipelineStage):
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1],
inference_args.num_frames - 1, batch.height,
fastvideo_args.num_frames - 1, batch.height,
batch.width)
],
dim=2)
@@ -70,23 +70,25 @@ class EncodingStage(PipelineStage):
dtype=torch.float32)
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not inference_args.disable_autocast
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
# Encode Image
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if inference_args.vae_tiling:
if fastvideo_args.vae_tiling:
self.vae.enable_tiling()
# if inference_args.vae_sp:
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
if not vae_autocast_enabled:
video_condition = video_condition.to(vae_dtype)
encoder_output = self.vae.encode(video_condition)
generator = batch.generator
if generator is None:
raise ValueError("Generator must be provided")
latent_condition = self.retrieve_latents(encoder_output, generator[0])
# Apply shifting if needed
@@ -104,9 +106,9 @@ class EncodingStage(PipelineStage):
else:
latent_condition = latent_condition * self.vae.scaling_factor
mask_lat_size = torch.ones(1, 1, inference_args.num_frames,
mask_lat_size = torch.ones(1, 1, fastvideo_args.num_frames,
latent_height, latent_width)
mask_lat_size[:, :, list(range(1, inference_args.num_frames))] = 0
mask_lat_size[:, :, list(range(1, fastvideo_args.num_frames))] = 0
first_frame_mask = mask_lat_size[:, :, 0:1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask,
@@ -5,7 +5,7 @@ Input validation stage for diffusion pipelines.
import torch
from fastvideo.v1.inference_args import InferenceArgs
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
@@ -22,10 +22,10 @@ class InputValidationStage(PipelineStage):
"""
def _generate_seeds(self, batch: ForwardBatch,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Generate seeds for the inference"""
seed = inference_args.seed
num_videos_per_prompt = inference_args.num_videos
seed = fastvideo_args.seed
num_videos_per_prompt = fastvideo_args.num_videos
seeds = [seed + i for i in range(num_videos_per_prompt)]
batch.seeds = seeds
@@ -37,19 +37,19 @@ class InputValidationStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Validate and prepare inputs.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The validated batch information.
"""
self._generate_seeds(batch, inference_args)
self._generate_seeds(batch, fastvideo_args)
# Ensure prompt is properly formatted
if batch.prompt is None and batch.prompt_embeds is None:
@@ -91,6 +91,6 @@ class InputValidationStage(PipelineStage):
# Set data type if not already set
if batch.data_type is None:
batch.data_type = inference_args.precision
batch.data_type = fastvideo_args.precision
return batch
@@ -4,7 +4,7 @@ Latent preparation stage for diffusion pipelines.
"""
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.v1.inference_args import InferenceArgs
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.pipelines.pipeline_batch_info import ForwardBatch
@@ -29,14 +29,14 @@ class LatentPreparationStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Prepare initial latent variables for the diffusion process.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with prepared latent variables.
@@ -44,7 +44,7 @@ class LatentPreparationStage(PipelineStage):
# Adjust video length based on VAE version if needed
if hasattr(self, 'adjust_video_length'):
batch = self.adjust_video_length(self.vae, batch, inference_args)
batch = self.adjust_video_length(self.vae, batch, fastvideo_args)
# Determine batch size
if isinstance(batch.prompt, list):
batch_size = len(batch.prompt)
@@ -69,13 +69,16 @@ class LatentPreparationStage(PipelineStage):
if height is None or width is None:
raise ValueError("Height and width must be provided")
assert fastvideo_args.num_channels_latents is not None
assert fastvideo_args.vae_scale_factor is not None
# Calculate latent shape
shape = (
batch_size,
inference_args.num_channels_latents,
fastvideo_args.num_channels_latents,
num_frames,
height // inference_args.vae_scale_factor,
width // inference_args.vae_scale_factor,
height // fastvideo_args.vae_scale_factor,
width // fastvideo_args.vae_scale_factor,
)
# Validate generator if it's a list
@@ -104,13 +107,13 @@ class LatentPreparationStage(PipelineStage):
return batch
def adjust_video_length(self, vae: ParallelTiledVAE, batch: ForwardBatch,
inference_args: InferenceArgs) -> ForwardBatch:
fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""
Adjust video length based on VAE version.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with adjusted video length.
@@ -10,7 +10,7 @@ from typing import TypedDict
import torch
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.inference_args import InferenceArgs
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
@@ -61,19 +61,19 @@ class LlamaEncodingStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode the prompt into text encoder hidden states.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded prompt embeddings.
"""
if inference_args.use_cpu_offload:
if fastvideo_args.use_cpu_offload:
self.text_encoder = self.text_encoder.to(batch.device)
text = prompt_template_video["template"].format(batch.prompt)
@@ -120,9 +120,10 @@ class LlamaEncodingStage(PipelineStage):
crop_start = prompt_template_video.get("crop_start", -1)
negative_last_hidden_state = negative_last_hidden_state[:,
crop_start:]
assert batch.negative_prompt_embeds is not None
batch.negative_prompt_embeds.append(negative_last_hidden_state)
if inference_args.use_cpu_offload:
if fastvideo_args.use_cpu_offload:
self.text_encoder.to('cpu')
torch.cuda.empty_cache()
+6 -5
View File
@@ -7,7 +7,7 @@ This module contains implementations of prompt encoding stages for diffusion pip
import torch
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.inference_args import InferenceArgs
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
@@ -38,19 +38,19 @@ class T5EncodingStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode the prompt into text encoder hidden states.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded prompt embeddings.
"""
if inference_args.use_cpu_offload:
if fastvideo_args.use_cpu_offload:
self.text_encoder = self.text_encoder.to(batch.device)
text = batch.prompt
@@ -106,9 +106,10 @@ class T5EncodingStage(PipelineStage):
for u in neg_prompt_embeds
],
dim=0)
assert batch.negative_prompt_embeds is not None
batch.negative_prompt_embeds.append(neg_prompt_embeds)
if inference_args.use_cpu_offload:
if fastvideo_args.use_cpu_offload:
self.text_encoder.to('cpu')
torch.cuda.empty_cache()
@@ -7,7 +7,7 @@ This module contains implementations of timestep preparation stages for diffusio
import inspect
from fastvideo.v1.inference_args import InferenceArgs
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
@@ -29,14 +29,14 @@ class TimestepPreparationStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Prepare timesteps for the diffusion process.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with prepared timesteps.
@@ -6,7 +6,7 @@ This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.v1.inference_args import InferenceArgs
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.stages import (
@@ -14,8 +14,6 @@ from fastvideo.v1.pipelines.stages import (
EncodingStage, InputValidationStage, LatentPreparationStage,
T5EncodingStage, TimestepPreparationStage)
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
@@ -26,7 +24,7 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
"image_encoder", "image_processor"
]
def create_pipeline_stages(self, inference_args: InferenceArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
@@ -67,15 +65,15 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def initialize_pipeline(self, inference_args: InferenceArgs):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
inference_args.vae_scale_factor = vae_scale_factor
fastvideo_args.vae_scale_factor = vae_scale_factor
num_channels_latents = self.get_module("transformer").out_channels
inference_args.num_channels_latents = num_channels_latents
fastvideo_args.num_channels_latents = num_channels_latents
EntryClass = WanImageToVideoPipeline
+5 -5
View File
@@ -6,7 +6,7 @@ This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.v1.inference_args import InferenceArgs
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.stages import (ConditioningStage, DecodingStage,
@@ -26,7 +26,7 @@ class WanPipeline(ComposedPipelineBase):
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def create_pipeline_stages(self, inference_args: InferenceArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
@@ -58,15 +58,15 @@ class WanPipeline(ComposedPipelineBase):
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def initialize_pipeline(self, inference_args: InferenceArgs):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
inference_args.vae_scale_factor = vae_scale_factor
fastvideo_args.vae_scale_factor = vae_scale_factor
num_channels_latents = self.get_module("transformer").in_channels
inference_args.num_channels_latents = num_channels_latents
fastvideo_args.num_channels_latents = num_channels_latents
EntryClass = WanPipeline
+1 -1
View File
@@ -18,7 +18,7 @@ def cuda_platform_plugin() -> Optional[str]:
try:
from fastvideo.v1.utils import import_pynvml
pynvml = import_pynvml()
pynvml = import_pynvml() # type: ignore[no-untyped-call]
pynvml.nvmlInit()
try:
# NOTE: Edge case: vllm cpu build on a GPU machine.
+7 -8
View File
@@ -26,7 +26,7 @@ logger = init_logger(__name__)
_P = ParamSpec("_P")
_R = TypeVar("_R")
pynvml = import_pynvml()
pynvml = import_pynvml() # type: ignore[no-untyped-call]
# pytorch 2.5 uses cudnn sdpa by default, which will cause crash on some models
# see https://github.com/huggingface/diffusers/issues/9704 for details
@@ -110,14 +110,13 @@ class CudaPlatformBase(Platform):
return float(torch.cuda.max_memory_allocated(device))
@classmethod
def get_attn_backend_cls(cls, selected_backend, head_size, dtype,
distributed) -> str:
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
head_size: int, dtype: torch.dtype) -> str:
# TODO(will): maybe come up with a more general interface for local attention
# if distributed is False, we always try to use Flash attn
logger.info(
"Distributed attention=%s, trying FASTVIDEO_ATTENTION_BACKEND=%s",
distributed, envs.FASTVIDEO_ATTENTION_BACKEND)
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
envs.FASTVIDEO_ATTENTION_BACKEND)
if selected_backend == _Backend.SLIDING_TILE_ATTN:
try:
from st_attn import sliding_tile_attention # noqa: F401
@@ -132,10 +131,10 @@ class CudaPlatformBase(Platform):
logger.info(
"Sliding Tile Attention backend is not installed. Fall back to Flash Attention."
)
elif selected_backend == _Backend.FLASH_ATTN:
pass
elif selected_backend == _Backend.TORCH_SDPA:
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
elif selected_backend == _Backend.FLASH_ATTN or selected_backend is None:
pass
elif selected_backend:
raise ValueError(f"Invalid attention backend for {cls.device_name}")
+2 -2
View File
@@ -87,8 +87,8 @@ class Platform:
return self._enum == PlatformEnum.CUDA
@classmethod
def get_attn_backend_cls(cls, selected_backend: _Backend, head_size: int,
dtype: torch.dtype, distributed: bool) -> str:
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
head_size: int, dtype: torch.dtype) -> str:
"""Get the attention backend class of a device."""
return ""
+22 -18
View File
@@ -11,12 +11,12 @@ from einops import rearrange
from fastvideo.v1.distributed import (init_distributed_environment,
initialize_model_parallel)
from fastvideo.v1.inference_args import InferenceArgs, prepare_inference_args
from fastvideo.v1.fastvideo_args import FastVideoArgs, prepare_fastvideo_args
# Fix the import path
from fastvideo.v1.inference_engine import InferenceEngine
def initialize_distributed_and_parallelism(inference_args: InferenceArgs):
def initialize_distributed_and_parallelism(fastvideo_args: FastVideoArgs):
local_rank = int(os.environ.get("LOCAL_RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
@@ -25,29 +25,33 @@ def initialize_distributed_and_parallelism(inference_args: InferenceArgs):
rank=rank,
local_rank=local_rank)
device_str = f"cuda:{local_rank}"
inference_args.device_str = device_str
inference_args.device = torch.device(device_str)
fastvideo_args.device_str = device_str
fastvideo_args.device = torch.device(device_str)
assert fastvideo_args.sp_size is not None
assert fastvideo_args.tp_size is not None
initialize_model_parallel(
sequence_model_parallel_size=inference_args.sp_size,
tensor_model_parallel_size=inference_args.tp_size,
sequence_model_parallel_size=fastvideo_args.sp_size,
tensor_model_parallel_size=fastvideo_args.tp_size,
)
def main(inference_args: InferenceArgs):
initialize_distributed_and_parallelism(inference_args)
engine = InferenceEngine.create_engine(inference_args, )
def main(fastvideo_args: FastVideoArgs):
initialize_distributed_and_parallelism(fastvideo_args)
engine = InferenceEngine.create_engine(fastvideo_args, )
if inference_args.prompt_path is not None:
with open(inference_args.prompt_path) as f:
if fastvideo_args.prompt_path is not None:
with open(fastvideo_args.prompt_path) as f:
prompts = [line.strip() for line in f.readlines()]
else:
prompts = [inference_args.prompt]
if fastvideo_args.prompt is None:
raise ValueError("prompt or prompt_path is required")
prompts = [fastvideo_args.prompt]
# Process each prompt
for prompt in prompts:
outputs = engine.run(
prompt=prompt,
inference_args=inference_args,
fastvideo_args=fastvideo_args,
)
# Process outputs
@@ -59,13 +63,13 @@ def main(inference_args: InferenceArgs):
frames.append((x * 255).numpy().astype(np.uint8))
# Save video
os.makedirs(os.path.dirname(inference_args.output_path), exist_ok=True)
imageio.mimsave(os.path.join(inference_args.output_path,
os.makedirs(os.path.dirname(fastvideo_args.output_path), exist_ok=True)
imageio.mimsave(os.path.join(fastvideo_args.output_path,
f"{prompt[:100]}.mp4"),
frames,
fps=inference_args.fps)
fps=fastvideo_args.fps)
if __name__ == "__main__":
inference_args = prepare_inference_args(sys.argv[1:])
main(inference_args)
fastvideo_args = prepare_fastvideo_args(sys.argv[1:])
main(fastvideo_args)
@@ -11,7 +11,7 @@ from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
load_tokenizer)
# from fastvideo.v1.models.hunyuan.text_encoder import load_text_encoder, load_tokenizer
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import maybe_download_model
@@ -38,7 +38,7 @@ def test_clip_encoder():
- Load models with the same weights and parameters
- Produce nearly identical outputs for the same input prompts
"""
args = InferenceArgs(model_path="openai/clip-vit-large-patch14",
args = FastVideoArgs(model_path="openai/clip-vit-large-patch14",
precision="float16")
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
@@ -9,7 +9,7 @@ from transformers import AutoConfig
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
load_tokenizer)
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
from fastvideo.v1.utils import maybe_download_model
@@ -38,7 +38,7 @@ def test_llama_encoder():
- Load models with the same weights and parameters
- Produce nearly identical outputs for the same input prompts
"""
args = InferenceArgs(model_path="meta-llama/Llama-2-7b-hf",
args = FastVideoArgs(model_path="meta-llama/Llama-2-7b-hf",
precision="float16")
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
@@ -7,7 +7,7 @@ import torch
from diffusers import WanTransformer3DModel
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import TransformerLoader
from fastvideo.v1.utils import maybe_download_model
@@ -29,7 +29,7 @@ def test_wan_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = InferenceArgs(model_path=TRANSFORMER_PATH,
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=False,
precision=precision_str)
args.device = device
+34 -24
View File
@@ -6,7 +6,7 @@ import pytest
import torch
from diffusers import AutoencoderKLWan
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader
from fastvideo.v1.utils import maybe_download_model
@@ -28,11 +28,12 @@ def test_wan_vae():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = InferenceArgs(model_path=VAE_PATH, vae_precision=precision_str)
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
args.device = device
loader = VAELoader()
model2 = loader.load(VAE_PATH, "", args)
assert model2.use_feature_cache # Default to use the original WanVAE algorithm
model1 = AutoencoderKLWan.from_pretrained(
VAE_PATH, torch_dtype=precision).to(device).eval()
@@ -48,43 +49,52 @@ def test_wan_vae():
32,
device=device,
dtype=precision)
latent_tensor = torch.randn(batch_size,
16,
21,
32,
32,
device=device,
dtype=precision)
# latent_tensor = torch.randn(batch_size,
# 16,
# 21,
# 32,
# 32,
# device=device,
# dtype=precision)
# Disable gradients for inference
with torch.no_grad():
# Test encoding
logger.info("Testing encoding...")
latent1 = model1.encode(input_tensor).latent_dist.mean
latent1 = model1.encode(input_tensor).latent_dist
print("--------------------------------")
latent2 = model2.encode(input_tensor).mean
latent2 = model2.encode(input_tensor)
# Check if latents have the same shape
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
assert latent1.mean.shape == latent2.mean.shape, f"Latent shapes don't match: {latent1.mean.shape} vs {latent2.mean.shape}"
# Check if latents are similar
max_diff_encode = torch.max(torch.abs(latent1 - latent2))
mean_diff_encode = torch.mean(torch.abs(latent1 - latent2))
max_diff_encode = torch.max(torch.abs(latent1.mean - latent2.mean))
mean_diff_encode = torch.mean(torch.abs(latent1.mean - latent2.mean))
logger.info("Maximum difference between encoded latents: %s",
max_diff_encode.item())
logger.info("Mean difference between encoded latents: %s",
mean_diff_encode.item())
assert mean_diff_encode < 5e-1, f"Encoded latents differ significantly: mean diff = {mean_diff_encode.item()}"
assert max_diff_encode < 1e-5, f"Encoded latents differ significantly: max diff = {mean_diff_encode.item()}"
# Test decoding
logger.info("Testing decoding...")
latent1_tensor = latent1.mode()
latents_mean = (torch.tensor(model1.config.latents_mean).view(
1, model1.config.z_dim, 1, 1, 1).to(latent_tensor.device,
latent_tensor.dtype))
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype))
latents_std = 1.0 / torch.tensor(model1.config.latents_std).view(
1, model1.config.z_dim, 1, 1, 1).to(latent_tensor.device,
latent_tensor.dtype)
latent_tensor = latent_tensor / latents_std + latents_mean
output2 = model2.decode(latent_tensor)
output1 = model1.decode(latent_tensor).sample
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype)
latent1_tensor = latent1_tensor / latents_std + latents_mean
output1 = model1.decode(latent1_tensor).sample
latent2_tensor = latent2.mode()
latents_mean = (torch.tensor(model2.config.latents_mean).view(
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype))
latents_std = 1.0 / torch.tensor(model2.config.latents_std).view(
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype)
latent2_tensor = latent2_tensor / latents_std + latents_mean
output2 = model2.decode(latent2_tensor)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
@@ -95,4 +105,4 @@ def test_wan_vae():
max_diff_decode.item())
logger.info("Mean difference between decoded outputs: %s",
mean_diff_decode.item())
assert mean_diff_decode < 1e-1, f"Decoded outputs differ significantly: mean diff = {mean_diff_decode.item()}"
assert max_diff_decode < 1e-5, f"Decoded outputs differ significantly: max diff = {mean_diff_decode.item()}"
+110 -11
View File
@@ -10,8 +10,10 @@ import math
import os
import sys
import tempfile
from functools import wraps
from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast
from functools import wraps, partial
from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast, Callable
from dataclasses import asdict, fields
import cloudpickle
import filelock
import torch
@@ -25,7 +27,7 @@ logger = init_logger(__name__)
T = TypeVar("T")
# TODO(will): used to convert inference_args.precision to torch.dtype. Find a
# TODO(will): used to convert fastvideo_args.precision to torch.dtype. Find a
# cleaner way to do this.
PRECISION_TO_TYPE = {
"fp32": torch.float32,
@@ -123,13 +125,14 @@ class SortedHelpFormatter(argparse.HelpFormatter):
class FlexibleArgumentParser(argparse.ArgumentParser):
"""ArgumentParser that allows both underscore and dash in names."""
def __init__(self, *args, **kwargs):
def __init__(self, *args, **kwargs) -> None:
# Set the default 'formatter_class' to SortedHelpFormatter
if 'formatter_class' not in kwargs:
kwargs['formatter_class'] = SortedHelpFormatter
super().__init__(*args, **kwargs)
def parse_args(self, args=None, namespace=None):
def parse_args( # type: ignore[override]
self, args=None, namespace=None) -> argparse.Namespace:
if args is None:
args = sys.argv[1:]
@@ -154,7 +157,8 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
else:
processed_args.append(arg)
return super().parse_args(processed_args, namespace)
return super().parse_args( # type: ignore[no-any-return]
processed_args, namespace)
def _pull_args_from_config(self, args: List[str]) -> List[str]:
"""Method to pull arguments specified in the config file
@@ -326,7 +330,7 @@ def warn_for_unimplemented_methods(cls: Type[T]) -> Type[T]:
return cls
def align_to(value, alignment):
def align_to(value: int, alignment: int) -> int:
"""align height, width according to alignment
Args:
@@ -362,9 +366,7 @@ def import_pynvml():
install FastVideo. It provides a Python module named `pynvml`.
- `pynvml` (https://pypi.org/project/pynvml/): An unofficial wrapper.
Prior to version 12.0, it also provides a Python module `pynvml`,
and therefore conflicts with the official one. What's worse,
the module is a Python package, and has higher priority than
the official one which is a standalone Python file.
and therefore conflicts with the official one which is a standalone Python file.
This causes errors when both of them are installed.
Starting from version 12.0, it migrates to a new module
named `pynvml_utils` to avoid the conflict.
@@ -381,12 +383,15 @@ def import_pynvml():
def maybe_download_model(model_path: str,
local_dir: Optional[str] = None) -> str:
local_dir: Optional[str] = None,
download: bool = True) -> str:
"""
Check if the model path is a Hugging Face Hub model ID and download it if needed.
Args:
model_path: Local path or Hugging Face Hub model ID
local_dir: Local directory to save the model
download: Whether to download the model from Hugging Face Hub
Returns:
Local path to the model
@@ -455,3 +460,97 @@ def verify_model_config_and_directory(model_path: str) -> Dict[str, Any]:
logger.info("Diffusers version: %s", config["_diffusers_version"])
return cast(Dict[str, Any], config)
def maybe_download_model_index(model_name_or_path: str) -> Dict[str, Any]:
"""
Download and extract just the model_index.json for a Hugging Face model.
Args:
model_name_or_path: Path or HF Hub model ID
Returns:
The parsed model_index.json as a dictionary
"""
import tempfile
from huggingface_hub import hf_hub_download
# If it's a local path, verify it directly
if os.path.exists(model_name_or_path):
return verify_model_config_and_directory(model_name_or_path)
# For remote models, download just the model_index.json
try:
with tempfile.TemporaryDirectory() as tmp_dir:
# Download just the model_index.json file
model_index_path = hf_hub_download(repo_id=model_name_or_path,
filename="model_index.json",
local_dir=tmp_dir)
# Load the model_index.json
with open(model_index_path) as f:
config: Dict[str, Any] = json.load(f)
# Verify it has the required fields
if "_class_name" not in config:
raise ValueError(
f"model_index.json for {model_name_or_path} does not contain _class_name field"
)
if "_diffusers_version" not in config:
raise ValueError(
f"model_index.json for {model_name_or_path} does not contain _diffusers_version field"
)
# Add the pipeline name for downstream use
config["pipeline_name"] = config["_class_name"]
logger.info("Downloaded model_index.json for %s, pipeline: %s",
model_name_or_path, config["_class_name"])
return config
except Exception as e:
raise ValueError(
f"Failed to download or parse model_index.json for {model_name_or_path}: {e}"
) from e
def update_environment_variables(envs: Dict[str, str]):
for k, v in envs.items():
if k in os.environ and os.environ[k] != v:
logger.warning(
"Overwriting environment variable %s "
"from '%s' to '%s'", k, os.environ[k], v)
os.environ[k] = v
def run_method(obj: Any, method: Union[str, bytes, Callable], args: tuple[Any],
kwargs: dict[str, Any]) -> Any:
"""
Run a method of an object with the given arguments and keyword arguments.
If the method is string, it will be converted to a method using getattr.
If the method is serialized bytes and will be deserialized using
cloudpickle.
If the method is a callable, it will be called directly.
"""
if isinstance(method, bytes):
func = partial(cloudpickle.loads(method), obj)
elif isinstance(method, str):
try:
func = getattr(obj, method)
except AttributeError:
raise NotImplementedError(f"Method {method!r} is not"
" implemented.") from None
else:
func = partial(method, obj) # type: ignore
return func(*args, **kwargs)
def diff_keys(a, b):
return [k for k in asdict(a) if asdict(a)[k] != asdict(b)[k]]
def update_in_place(target, source, ignore_fields=()):
for f in fields(target):
if hasattr(source, f.name) and f.name not in list(ignore_fields):
setattr(target, f.name, getattr(source, f.name))
+2 -2
View File
@@ -36,8 +36,8 @@ dependencies = [
"wandb==0.18.5", "loguru", "test-tube==0.7.5",
# Miscellaneous Utilities
"tqdm==4.66.5", "PyYAML==6.0.1", "idna==3.6", "protobuf==5.28.3",
"gradio==5.3.0", "moviepy==1.0.3", "flask",
"tqdm==4.66.5", "PyYAML==6.0.1", "protobuf==5.28.3",
"gradio>=5.22.0", "moviepy==1.0.3", "flask",
"flask_restful", "aiohttp", "huggingface_hub", "cloudpickle",
# System & Monitoring Tools
"gpustat", "watch",
+6 -6
View File
@@ -1,23 +1,23 @@
#!/bin/bash
num_gpus=4
export FASTVIDEO_ATTENTION_BACKEND=
export MODEL_BASE=FastVideo/FastHunyuan-diffusers
export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--sp_size 4 \
--tp_size 4 \
--sp_size $num_gpus \
--tp_size $num_gpus \
--height 720 \
--width 1280 \
--num_frames 125 \
--num_inference_steps 6 \
--num_inference_steps 50 \
--guidance_scale 1 \
--embedded_cfg_scale 6 \
--flow_shift 17 \
--flow_shift 7 \
--prompt_path ./assets/prompt.txt \
--seed 1024 \
--output_path outputs_video/ \
@@ -1,24 +1,24 @@
#!/bin/bash
num_gpus=1
num_gpus=2
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_hunyuan.json
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export MODEL_BASE=FastVideo/FastHunyuan-diffusers
export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--sp_size 1 \
--tp_size 1 \
--sp_size ${num_gpus} \
--tp_size ${num_gpus} \
--height 768 \
--width 1280 \
--num_frames 117 \
--num_inference_steps 6 \
--num_inference_steps 50 \
--guidance_scale 1 \
--embedded_cfg_scale 6 \
--flow_shift 17 \
--flow_shift 7 \
--prompt_path ./assets/prompt.txt \
--seed 1024 \
--output_path outputs_video/ \
+1 -1
View File
@@ -2,7 +2,7 @@
num_gpus=2
export FASTVIDEO_ATTENTION_BACKEND=
export MODEL_BASE=/workspace/data/Wan2.1-T2V-1.3B-Diffusers
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
+1 -1
View File
@@ -2,7 +2,7 @@
num_gpus=2
export FASTVIDEO_ATTENTION_BACKEND=
export MODEL_BASE=/workspace/data/Wan2.1-I2V-14B-480P-Diffusers
export MODEL_BASE=Wan-AI/Wan2.1-I2V-14B-480P-Diffusers
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for