Compare commits
16
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6b5b2dc6e7 | ||
|
|
91b7cc1be8 | ||
|
|
d6365373b4 | ||
|
|
50fb94b902 | ||
|
|
2caa0d4d0b | ||
|
|
24db823998 | ||
|
|
8c4704edf5 | ||
|
|
dfba7ec833 | ||
|
|
7a2e171f1b | ||
|
|
a8aac6090a | ||
|
|
191d1be3b4 | ||
|
|
338ea1e5f2 | ||
|
|
42a2f272d5 | ||
|
|
982bfcfdc8 | ||
|
|
2ac06379a7 | ||
|
|
baaa1673f7 |
@@ -4,6 +4,14 @@ title: "[Bug] "
|
||||
labels: ['Bug']
|
||||
|
||||
body:
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/env_utils.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Describe the bug
|
||||
@@ -17,13 +25,5 @@ body:
|
||||
What command or script did you run? Which **model** are you using?
|
||||
placeholder: |
|
||||
A placeholder for the command.
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
@@ -39,16 +39,6 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_training_test:
|
||||
description: "Run training-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_nightly_test:
|
||||
description: "Run nightly-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
@@ -69,7 +59,6 @@ jobs:
|
||||
encoder-test: ${{ steps.filter.outputs.encoder-test }}
|
||||
vae-test: ${{ steps.filter.outputs.vae-test }}
|
||||
transformer-test: ${{ steps.filter.outputs.transformer-test }}
|
||||
training-test: ${{ steps.filter.outputs.training-test }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
@@ -88,10 +77,6 @@ jobs:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
training-test:
|
||||
- 'fastvideo/v1/**'
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -156,8 +141,8 @@ jobs:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
python-version: [
|
||||
# {version: "3.10", tag: "latest"},
|
||||
# {version: "3.11", tag: "py3.11-latest"},
|
||||
{version: "3.10", tag: "latest"},
|
||||
{version: "3.11", tag: "py3.11-latest"},
|
||||
{version: "3.12", tag: "py3.12-latest"}
|
||||
]
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
@@ -173,43 +158,6 @@ jobs:
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
training-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "training-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/training -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
nightly-test:
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "nightly-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
runpod-cleanup:
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
|
||||
|
||||
@@ -10,7 +10,7 @@ jobs:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
python-version: "3.10"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
|
||||
@@ -43,8 +43,6 @@ on:
|
||||
required: true
|
||||
RUNPOD_PRIVATE_KEY:
|
||||
required: true
|
||||
WANDB_API_KEY:
|
||||
required: false
|
||||
|
||||
jobs:
|
||||
run-test:
|
||||
@@ -57,7 +55,7 @@ jobs:
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
@@ -74,7 +72,6 @@ jobs:
|
||||
JOB_ID: ${{ inputs.job_id }}
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
timeout-minutes: ${{ inputs.timeout_minutes }}
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
|
||||
@@ -5,7 +5,7 @@ on:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/sliding_tile_attention/setup.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
@@ -23,13 +23,13 @@ jobs:
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/attn
|
||||
cd csrc/sliding_tile_attention
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
@@ -136,21 +136,19 @@ jobs:
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py bdist_wheel --dist-dir=dist
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/attn
|
||||
cd csrc/sliding_tile_attention
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
|
||||
@@ -165,7 +163,7 @@ jobs:
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/attn/dist/*.whl
|
||||
path: csrc/sliding_tile_attention/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
@@ -231,19 +229,17 @@ jobs:
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py sdist --dist-dir=dist
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/attn/dist/
|
||||
packages-dir: csrc/sliding_tile_attention/dist/
|
||||
|
||||
@@ -28,4 +28,4 @@ jobs:
|
||||
|
||||
- name: Run Pytest
|
||||
run: |
|
||||
pytest --ignore csrc/attn/test
|
||||
pytest --ignore csrc/sliding_tile_attention/test
|
||||
|
||||
@@ -27,6 +27,7 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
**.json
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
|
||||
+2
-2
@@ -1,3 +1,3 @@
|
||||
[submodule "csrc/attn/tk"]
|
||||
path = csrc/attn/tk
|
||||
[submodule "csrc/sliding_tile_attention/tk"]
|
||||
path = csrc/sliding_tile_attention/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
@@ -33,7 +33,7 @@ repos:
|
||||
args: [--in-place, --verbose]
|
||||
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.11.12
|
||||
rev: v0.11.4
|
||||
hooks:
|
||||
- id: ruff
|
||||
args: [--output-format, github, --fix]
|
||||
@@ -48,7 +48,7 @@ repos:
|
||||
hooks:
|
||||
- id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.30
|
||||
rev: v0.9.29
|
||||
hooks:
|
||||
- id: pymarkdown
|
||||
args: [fix]
|
||||
|
||||
+42491
-42491
File diff suppressed because it is too large
Load Diff
@@ -70,8 +70,6 @@ DEFAULT_CONDA_PATTERNS = {
|
||||
"optree",
|
||||
"nccl",
|
||||
"transformers",
|
||||
"accelerate",
|
||||
"peft",
|
||||
"zmq",
|
||||
"nvidia",
|
||||
"pynvml",
|
||||
@@ -87,8 +85,6 @@ DEFAULT_PIP_PATTERNS = {
|
||||
"onnx",
|
||||
"nccl",
|
||||
"transformers",
|
||||
"accelerate",
|
||||
"peft",
|
||||
"zmq",
|
||||
"nvidia",
|
||||
"pynvml",
|
||||
@@ -1,225 +0,0 @@
|
||||
import torch
|
||||
import argparse
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
def set_seed(seed: int = 42):
|
||||
# Python random module
|
||||
random.seed(seed)
|
||||
|
||||
# NumPy
|
||||
np.random.seed(seed)
|
||||
|
||||
# PyTorch
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed) # if using multi-GPU
|
||||
|
||||
def parse_arguments():
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=1, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=64, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
|
||||
return parser.parse_args()
|
||||
|
||||
def create_input_tensors(batch, head, seq_len, headdim):
|
||||
"""Create random input tensors for attention."""
|
||||
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
return q, k, v
|
||||
|
||||
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
|
||||
|
||||
Args:
|
||||
bs: batch size
|
||||
h: number of heads
|
||||
num_q_blocks: number of query blocks
|
||||
num_kv_blocks: number of key-value blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to k).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Binary mask where 1 indicates attention connection.
|
||||
"""
|
||||
# Ensure k is not larger than num_kv_blocks
|
||||
k = min(k, num_kv_blocks)
|
||||
|
||||
# Create random scores for sampling
|
||||
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
|
||||
|
||||
# Get top-k indices for each q block
|
||||
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
|
||||
|
||||
# sort q2k_block_sparse_index
|
||||
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
|
||||
|
||||
# All q blocks attend to exactly k kv blocks
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
|
||||
|
||||
# Create the corresponding mask
|
||||
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
# Fill in the mask based on the indices
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx]
|
||||
block_sparse_mask[b, head, q_idx, kv_indices] = True
|
||||
|
||||
# Create the reverse mapping (k2q)
|
||||
# First, initialize lists to collect q indices for each kv block
|
||||
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
|
||||
|
||||
# Populate the lists based on q2k mapping
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
|
||||
for kv_idx in kv_indices:
|
||||
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
|
||||
|
||||
# Find the maximum number of q blocks that attend to any kv block
|
||||
max_q_per_kv = 0
|
||||
for flat_idx in range(bs * h):
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
|
||||
|
||||
# Create tensors for k2q mapping
|
||||
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
|
||||
dtype=torch.int32, device=device)
|
||||
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
|
||||
dtype=torch.int32, device=device)
|
||||
|
||||
# Fill the tensors
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
q_indices = k2q_indices_list[flat_idx][kv_idx]
|
||||
num_q = len(q_indices)
|
||||
k2q_block_sparse_num[b, head, kv_idx] = num_q
|
||||
if num_q > 0:
|
||||
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
|
||||
q_indices, dtype=torch.int32, device=device)
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
|
||||
|
||||
def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops):
|
||||
"""Benchmark block sparse attention forward and backward passes."""
|
||||
print("\n=== BLOCK SPARSE ATTENTION BENCHMARK ===")
|
||||
|
||||
# Forward pass
|
||||
# Warm-up run
|
||||
o, l_vec = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark forward
|
||||
_, fwd_time = benchmark_forward(
|
||||
block_sparse_attention_fwd,
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num,
|
||||
repeats=20,
|
||||
verbose=False,
|
||||
desc='Block Sparse Forward'
|
||||
)
|
||||
|
||||
sparse_tflops = flops / fwd_time.mean * 1e-12
|
||||
print(f"Block Sparse Forward - TFLOPS: {sparse_tflops:.2f}")
|
||||
|
||||
# Backward pass
|
||||
grad_output = torch.randn_like(o)
|
||||
|
||||
# Warm-up runs
|
||||
for _ in range(5):
|
||||
block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark backward
|
||||
_, bwd_time = benchmark_forward(
|
||||
block_sparse_attention_backward,
|
||||
q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num,
|
||||
repeats=20,
|
||||
verbose=False,
|
||||
desc='Block Sparse Backward'
|
||||
)
|
||||
bwd_flops = 2.5 * flops # Approximation
|
||||
|
||||
sparse_bwd_tflops = bwd_flops / bwd_time.mean * 1e-12
|
||||
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd_tflops:.2f}")
|
||||
|
||||
return sparse_tflops, sparse_bwd_tflops
|
||||
|
||||
def main():
|
||||
args = parse_arguments()
|
||||
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
batch = args.batch_size
|
||||
head = args.num_heads
|
||||
headdim = args.head_dim
|
||||
|
||||
print(f"Block Sparse Attention Benchmark")
|
||||
print(f"batch: {batch}, head: {head}, headdim: {headdim}")
|
||||
|
||||
# Test with different sequence lengths
|
||||
for seq_len in args.seq_lengths:
|
||||
# Skip very long sequences if they might cause OOM
|
||||
if seq_len > 16384 and batch > 1:
|
||||
continue
|
||||
|
||||
print("="*100)
|
||||
print(f"\nSequence length: {seq_len}")
|
||||
|
||||
# Calculate theoretical FLOPs for attention
|
||||
flops = 4 * batch * head * headdim * seq_len * seq_len
|
||||
|
||||
# Create input tensors
|
||||
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
|
||||
|
||||
# Setup block sparse parameters
|
||||
num_q_blocks = seq_len // BLOCK_M
|
||||
num_kv_blocks = seq_len // BLOCK_N
|
||||
|
||||
# Determine k value (number of kv blocks per q block)
|
||||
topk = args.topk
|
||||
if topk is None:
|
||||
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
|
||||
topk = max(1, topk)
|
||||
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
|
||||
|
||||
# Generate block sparse pattern
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, _ = generate_block_sparse_pattern(
|
||||
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
|
||||
|
||||
# Benchmark block sparse attention
|
||||
sparse_fwd, sparse_bwd = benchmark_block_sparse_attention(
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops
|
||||
)
|
||||
|
||||
# Print results
|
||||
print("\n=== PERFORMANCE RESULTS ===")
|
||||
print(f"Block Sparse Forward - TFLOPS: {sparse_fwd:.2f}")
|
||||
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd:.2f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,15 +0,0 @@
|
||||
### ADD TO THIS TO REGISTER NEW KERNELS
|
||||
sources = {
|
||||
'block_sparse': {
|
||||
'source_files': {
|
||||
'h100': 'vsa/block_sparse_h100.cu'
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
### WHICH KERNELS DO WE WANT TO BUILD?
|
||||
# (oftentimes during development work you don't need to redefine them all.)
|
||||
kernels = ['block_sparse']
|
||||
|
||||
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
|
||||
target = 'h100'
|
||||
@@ -1,76 +0,0 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from csrc.attn.config_vsa import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "vsa"
|
||||
VERSION = "0.0.1"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn"
|
||||
|
||||
# Set environment variables
|
||||
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
|
||||
python_include = subprocess.check_output(['python', '-c',
|
||||
"import sysconfig; print(sysconfig.get_path('include'))"]).decode().strip()
|
||||
torch_include = subprocess.check_output([
|
||||
'python', '-c',
|
||||
"import torch; from torch.utils.cpp_extension import include_paths; print(' '.join(['-I' + p for p in include_paths()]))"
|
||||
]).decode().strip()
|
||||
print('vsa root:', tk_root)
|
||||
print('Python include:', python_include)
|
||||
print('Torch include directories:', torch_include)
|
||||
|
||||
# CUDA flags
|
||||
cuda_flags = [
|
||||
'-DNDEBUG', '-Xcompiler=-Wno-psabi', '-Xcompiler=-fno-strict-aliasing', '--expt-extended-lambda',
|
||||
'--expt-relaxed-constexpr', '-forward-unknown-to-host-compiler', '--use_fast_math', '-std=c++20', '-O3',
|
||||
'-Xnvlink=--verbose', '-Xptxas=--verbose', '-Xptxas=--warn-on-spills', f'-I{tk_root}/include',
|
||||
f'-I{tk_root}/prototype', f'-I{python_include}', '-DTORCH_COMPILE'
|
||||
] + torch_include.split()
|
||||
cpp_flags = ['-std=c++20', '-O3']
|
||||
|
||||
if target == 'h100':
|
||||
cuda_flags.append('-DKITTENS_HOPPER')
|
||||
cuda_flags.append('-arch=sm_90a')
|
||||
else:
|
||||
raise ValueError(f'Target {target} not supported')
|
||||
|
||||
source_files = ['vsa.cpp']
|
||||
for k in kernels:
|
||||
if target not in sources[k]['source_files']:
|
||||
raise KeyError(f'Target {target} not found in source files for kernel {k}')
|
||||
if isinstance(sources[k]['source_files'][target], list):
|
||||
source_files.extend(sources[k]['source_files'][target])
|
||||
else:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=[
|
||||
CUDAExtension('vsa_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
],
|
||||
cmdclass={'build_ext': BuildExtension},
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
"Environment :: GPU :: NVIDIA CUDA :: 12",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.10',
|
||||
install_requires=["torch>=2.5.0"])
|
||||
@@ -1,266 +0,0 @@
|
||||
import torch
|
||||
import argparse
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
from flash_attn import flash_attn_func
|
||||
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward, BlockSparseAttentionFunction
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
def set_seed(seed: int = 42):
|
||||
# Python random module
|
||||
random.seed(seed)
|
||||
|
||||
# NumPy
|
||||
np.random.seed(seed)
|
||||
|
||||
# PyTorch
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed) # if using multi-GPU
|
||||
|
||||
def parse_arguments():
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
|
||||
parser.add_argument('--num_iterations', type=int, default=100, help='Number of test iterations to run')
|
||||
return parser.parse_args()
|
||||
|
||||
@torch.no_grad
|
||||
def precision_metric(quant_o, fa2_o):
|
||||
x, xx = quant_o.float(), fa2_o.float()
|
||||
sim = torch.nn.functional.cosine_similarity(x.reshape(1, -1), xx.reshape(1, -1)).item()
|
||||
l1 = ((x - xx).abs().sum() / xx.abs().sum() ).item()
|
||||
rmse = torch.sqrt(torch.mean((x -xx) ** 2)).item()
|
||||
|
||||
return sim, l1, rmse
|
||||
|
||||
def create_input_tensors(batch, head, seq_len, headdim):
|
||||
"""Create random input tensors for attention."""
|
||||
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
return q, k, v
|
||||
|
||||
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
|
||||
|
||||
Args:
|
||||
bs: batch size
|
||||
h: number of heads
|
||||
num_q_blocks: number of query blocks
|
||||
num_kv_blocks: number of key-value blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to k).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Binary mask where 1 indicates attention connection.
|
||||
"""
|
||||
# Ensure k is not larger than num_kv_blocks
|
||||
k = min(k, num_kv_blocks)
|
||||
|
||||
# Create random scores for sampling
|
||||
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
|
||||
|
||||
# Get top-k indices for each q block
|
||||
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
|
||||
|
||||
# sort q2k_block_sparse_index
|
||||
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
|
||||
|
||||
# All q blocks attend to exactly k kv blocks
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
|
||||
|
||||
# Create the corresponding mask
|
||||
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
# Fill in the mask based on the indices
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx]
|
||||
block_sparse_mask[b, head, q_idx, kv_indices] = True
|
||||
|
||||
# Create the reverse mapping (k2q)
|
||||
# First, initialize lists to collect q indices for each kv block
|
||||
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
|
||||
|
||||
# Populate the lists based on q2k mapping
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
|
||||
for kv_idx in kv_indices:
|
||||
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
|
||||
|
||||
# Find the maximum number of q blocks that attend to any kv block
|
||||
max_q_per_kv = 0
|
||||
for flat_idx in range(bs * h):
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
|
||||
|
||||
# Create tensors for k2q mapping
|
||||
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
|
||||
dtype=torch.int32, device=device)
|
||||
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
|
||||
dtype=torch.int32, device=device)
|
||||
|
||||
# Fill the tensors
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
q_indices = k2q_indices_list[flat_idx][kv_idx]
|
||||
num_q = len(q_indices)
|
||||
k2q_block_sparse_num[b, head, kv_idx] = num_q
|
||||
if num_q > 0:
|
||||
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
|
||||
q_indices, dtype=torch.int32, device=device)
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
|
||||
|
||||
def main():
|
||||
args = parse_arguments()
|
||||
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
batch = args.batch_size
|
||||
head = args.num_heads
|
||||
headdim = args.head_dim
|
||||
num_iterations = args.num_iterations
|
||||
|
||||
print(f"Block Sparse Attention Benchmark")
|
||||
print(f"batch: {batch}, head: {head}, headdim: {headdim}, iterations: {num_iterations}")
|
||||
|
||||
# Test with different sequence lengths
|
||||
for seq_len in args.seq_lengths:
|
||||
# Skip very long sequences if they might cause OOM
|
||||
# if seq_len > 16384 and batch > 1:
|
||||
# continue
|
||||
|
||||
print("="*100)
|
||||
print(f"\nSequence length: {seq_len}")
|
||||
|
||||
# Collect metrics across iterations
|
||||
forward_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_q_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_k_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_v_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
|
||||
for iter_idx in range(num_iterations):
|
||||
if num_iterations > 1:
|
||||
print(f"\nIteration {iter_idx+1}/{num_iterations}")
|
||||
|
||||
# Create input tensors
|
||||
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
|
||||
|
||||
# Setup block sparse parameters
|
||||
num_q_blocks = seq_len // BLOCK_M
|
||||
num_kv_blocks = seq_len // BLOCK_N
|
||||
|
||||
# Determine k value (number of kv blocks per q block)
|
||||
topk = args.topk
|
||||
if topk is None:
|
||||
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
|
||||
topk = max(1, topk)
|
||||
if iter_idx == 0: # Only print this once
|
||||
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
|
||||
|
||||
# Generate block sparse pattern
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask = generate_block_sparse_pattern(
|
||||
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
|
||||
|
||||
# expand block_sparse_mask to full mask
|
||||
block_mask_expanded = block_sparse_mask.unsqueeze(-1).unsqueeze(-2) # [b, h, num_q_blocks, num_kv_blocks, 1, 1]
|
||||
block_mask_expanded = block_mask_expanded.expand(-1, -1, -1, -1, BLOCK_M, BLOCK_N) # [b, h, num_q_blocks, num_kv_blocks, BLOCK_M, BLOCK_N]
|
||||
full_mask = block_mask_expanded.permute(0, 1, 2, 4, 3, 5).reshape(batch, head, seq_len, seq_len)
|
||||
|
||||
q_sdpa = q.clone()
|
||||
k_sdpa = k.clone()
|
||||
v_sdpa = v.clone()
|
||||
|
||||
q.requires_grad = True
|
||||
k.requires_grad = True
|
||||
v.requires_grad = True
|
||||
q_sdpa.requires_grad = True
|
||||
k_sdpa.requires_grad = True
|
||||
v_sdpa.requires_grad = True
|
||||
|
||||
# testing forward
|
||||
o = BlockSparseAttentionFunction.apply(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
o_sdpa = torch.nn.functional.scaled_dot_product_attention(q_sdpa, k_sdpa, v_sdpa, attn_mask=full_mask)
|
||||
|
||||
sim, l1, rmse = precision_metric(o, o_sdpa)
|
||||
forward_metrics['sim'].append(sim)
|
||||
forward_metrics['l1'].append(l1)
|
||||
forward_metrics['rmse'].append(rmse)
|
||||
|
||||
print(f"block_sparse_attention_fwd vs torch.nn.functional.scaled_dot_product_attention:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
# test backward
|
||||
grad_o = torch.randn_like(o)
|
||||
o.backward(grad_o)
|
||||
o_sdpa.backward(grad_o)
|
||||
|
||||
sim, l1, rmse = precision_metric(q.grad, q_sdpa.grad)
|
||||
grad_q_metrics['sim'].append(sim)
|
||||
grad_q_metrics['l1'].append(l1)
|
||||
grad_q_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_q:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
sim, l1, rmse = precision_metric(k.grad, k_sdpa.grad)
|
||||
grad_k_metrics['sim'].append(sim)
|
||||
grad_k_metrics['l1'].append(l1)
|
||||
grad_k_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_k:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
sim, l1, rmse = precision_metric(v.grad, v_sdpa.grad)
|
||||
grad_v_metrics['sim'].append(sim)
|
||||
grad_v_metrics['l1'].append(l1)
|
||||
grad_v_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_v:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
# Print summary statistics if multiple iterations were run
|
||||
if num_iterations > 1:
|
||||
print("\n" + "="*50)
|
||||
print(f"Summary Statistics (over {num_iterations} iterations):")
|
||||
|
||||
print("\nForward metrics:")
|
||||
print(f"Similarity: mean={np.mean(forward_metrics['sim']):.6f}, std={np.std(forward_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(forward_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient Q metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_q_metrics['sim']):.6f}, std={np.std(grad_q_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(grad_q_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient K metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_k_metrics['sim']):.6f}, std={np.std(grad_k_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(grad_k_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient V metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_v_metrics['sim']):.6f}, std={np.std(grad_v_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,136 +0,0 @@
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
|
||||
def pytorch_test(Q, K, V, dO):
|
||||
q_ = Q.to(torch.float64).requires_grad_()
|
||||
k_ = K.to(torch.float64).requires_grad_()
|
||||
v_ = V.to(torch.float64).requires_grad_()
|
||||
dO_ = dO.to(torch.float64)
|
||||
|
||||
# manual pytorch implementation of scaled dot product attention
|
||||
QK = torch.matmul(q_, k_.transpose(-2, -1))
|
||||
QK /= (q_.size(-1) ** 0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v_)
|
||||
|
||||
output.backward(dO_)
|
||||
|
||||
q_grad = q_.grad
|
||||
k_grad = k_.grad
|
||||
v_grad = v_.grad
|
||||
|
||||
return output, q_grad, k_grad, v_grad
|
||||
|
||||
def fa2_test(Q, K, V, dO):
|
||||
Q.requires_grad = True
|
||||
K.requires_grad = True
|
||||
V.requires_grad = True
|
||||
output = torch.nn.functional.scaled_dot_product_attention(Q, K, V, is_causal=False)
|
||||
output.backward(dO)
|
||||
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
def check_correctness(b, h, n, d, mean, std, num_iterations=100, error_mode='all', test_mode='forward_backward'):
|
||||
results = {
|
||||
'FA2 vs PT': {'sum_diff': 0, 'sum_abs': 0, 'max_diff': 0},
|
||||
}
|
||||
|
||||
for _ in range(num_iterations):
|
||||
torch.manual_seed(0)
|
||||
|
||||
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
dO = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
|
||||
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, dO)
|
||||
fa2_o, fa2_qg, fa2_kg, fa2_vg = fa2_test(Q, K, V, dO)
|
||||
|
||||
if test_mode == 'forward_only':
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
else: # 'forward_backward'
|
||||
if error_mode == 'output':
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
elif error_mode == 'backward':
|
||||
tensors_fa2_pt = [(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
else: # 'all'
|
||||
tensors_fa2_pt = [(pt_o, fa2_o),
|
||||
(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
|
||||
for pt, fa2 in tensors_fa2_pt:
|
||||
diff = pt - fa2
|
||||
abs_diff = torch.abs(diff)
|
||||
results['FA2 vs PT']['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results['FA2 vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
results['FA2 vs PT']['max_diff'] = max(results['FA2 vs PT']['max_diff'], torch.max(abs_diff).item())
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Calculate total elements based on test mode and error mode
|
||||
if test_mode == 'forward_only':
|
||||
total_elements = b * h * n * d * num_iterations
|
||||
else: # 'forward_backward'
|
||||
total_elements = b * h * n * d * num_iterations * (1 if error_mode == 'output' else 3 if error_mode == 'backward' else 4)
|
||||
|
||||
for name, data in results.items():
|
||||
avg_diff = data['sum_diff'] / total_elements
|
||||
max_diff = data['max_diff']
|
||||
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
|
||||
|
||||
return results
|
||||
|
||||
def generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward'):
|
||||
seq_lengths = [768 * (2**i) for i in range(1)]
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print(f"ATTENTION ERROR COMPARISON TABLE (b={b}, h={h}, d={d}, mean={mean}, std={std})")
|
||||
print(f"Mode: {error_mode}, Test: {test_mode}")
|
||||
print(f"{'='*80}")
|
||||
|
||||
# Print header
|
||||
print(f"{'Seq Length':<12} | {'FA2 vs PT Avg':<15} | {'FA2 vs PT Max':<15}")
|
||||
print(f"{'-'*12} | {'-'*15} | {'-'*15}")
|
||||
|
||||
for n in seq_lengths:
|
||||
results = check_correctness(b, h, n, d, mean, std, error_mode=error_mode, test_mode=test_mode)
|
||||
|
||||
fa2_pt_avg = results['FA2 vs PT']['avg_diff']
|
||||
fa2_pt_max = results['FA2 vs PT']['max_diff']
|
||||
|
||||
# Print row
|
||||
print(f"{n:<12} | {fa2_pt_avg:<15.6e} | {fa2_pt_max:<15.6e}")
|
||||
|
||||
print(f"{'='*80}\n")
|
||||
|
||||
# fix random seed
|
||||
torch.manual_seed(0)
|
||||
|
||||
# Example usage
|
||||
b, h, d = 2, 2, 64
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
# Test forward only
|
||||
generate_error_tables(b, h, d, mean, std, error_mode='output', test_mode='forward_only')
|
||||
|
||||
# Test forward and backward
|
||||
generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward')
|
||||
|
||||
print("Attention error comparison completed.")
|
||||
@@ -1,175 +0,0 @@
|
||||
import torch
|
||||
from flash_attn_interface import flash_attn_func
|
||||
from st_attn import mha_forward, mha_backward
|
||||
import random
|
||||
from tqdm import tqdm
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
|
||||
def pytorch_test(Q, K, V, dO):
|
||||
q_ = Q.to(torch.float64).requires_grad_()
|
||||
k_ = K.to(torch.float64).requires_grad_()
|
||||
v_ = V.to(torch.float64).requires_grad_()
|
||||
dO_ = dO.to(torch.float64)
|
||||
|
||||
# manual pytorch implementation of scaled dot product attention
|
||||
QK = torch.matmul(q_, k_.transpose(-2, -1))
|
||||
QK /= (q_.size(-1) ** 0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v_)
|
||||
|
||||
output.backward(dO_)
|
||||
|
||||
q_grad = q_.grad
|
||||
k_grad = k_.grad
|
||||
v_grad = v_.grad
|
||||
|
||||
return output, q_grad, k_grad, v_grad
|
||||
|
||||
def fa2_test(Q, K, V, dO):
|
||||
Q.requires_grad = True
|
||||
K.requires_grad = True
|
||||
V.requires_grad = True
|
||||
output = torch.nn.functional.scaled_dot_product_attention(Q, K, V, is_causal=False)
|
||||
output.backward(dO)
|
||||
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
|
||||
def mha_kernel_test(Q, K, V, dO, mode):
|
||||
Q.requires_grad = True
|
||||
K.requires_grad = True
|
||||
V.requires_grad = True
|
||||
|
||||
o, l_vec = mha_forward(Q, K, V)
|
||||
|
||||
if mode == 'forward_only':
|
||||
return o, None, None, None
|
||||
else: # 'forward_backward'
|
||||
qg, kg, vg = mha_backward(Q, K, V, o, l_vec, dO)
|
||||
return o, qg, kg, vg
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
def check_correctness(b, h, n, d, mean, std, num_iterations=100, error_mode='all', test_mode='forward_backward'):
|
||||
results = {
|
||||
'MHA vs PT': {'sum_diff': 0, 'sum_abs': 0, 'max_diff': 0},
|
||||
'FA2 vs PT': {'sum_diff': 0, 'sum_abs': 0, 'max_diff': 0},
|
||||
}
|
||||
|
||||
for _ in range(num_iterations):
|
||||
torch.manual_seed(0)
|
||||
|
||||
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
dO = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
|
||||
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, dO)
|
||||
fa2_o, fa2_qg, fa2_kg, fa2_vg = fa2_test(Q, K, V, dO)
|
||||
|
||||
if test_mode == 'forward_only':
|
||||
mha_o, _, _, _ = mha_kernel_test(Q, K, V, dO, 'forward_only')
|
||||
tensors_mha_pt = [(pt_o, mha_o)]
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
else: # 'forward_backward'
|
||||
mha_o, mha_qg, mha_kg, mha_vg = mha_kernel_test(Q, K, V, dO, 'forward_backward')
|
||||
|
||||
if error_mode == 'output':
|
||||
tensors_mha_pt = [(pt_o, mha_o)]
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
elif error_mode == 'backward':
|
||||
tensors_mha_pt = [(pt_qg, mha_qg),
|
||||
(pt_kg, mha_kg),
|
||||
(pt_vg, mha_vg)]
|
||||
tensors_fa2_pt = [(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
else: # 'all'
|
||||
tensors_mha_pt = [(pt_o, mha_o),
|
||||
(pt_qg, mha_qg),
|
||||
(pt_kg, mha_kg),
|
||||
(pt_vg, mha_vg)]
|
||||
tensors_fa2_pt = [(pt_o, fa2_o),
|
||||
(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
|
||||
for pt, mha in tensors_mha_pt:
|
||||
diff = pt - mha
|
||||
abs_diff = torch.abs(diff)
|
||||
results['MHA vs PT']['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results['MHA vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
results['MHA vs PT']['max_diff'] = max(results['MHA vs PT']['max_diff'], torch.max(abs_diff).item())
|
||||
|
||||
for pt, fa2 in tensors_fa2_pt:
|
||||
diff = pt - fa2
|
||||
abs_diff = torch.abs(diff)
|
||||
results['FA2 vs PT']['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results['FA2 vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
results['FA2 vs PT']['max_diff'] = max(results['FA2 vs PT']['max_diff'], torch.max(abs_diff).item())
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Calculate total elements based on test mode and error mode
|
||||
if test_mode == 'forward_only':
|
||||
total_elements = b * h * n * d * num_iterations
|
||||
else: # 'forward_backward'
|
||||
total_elements = b * h * n * d * num_iterations * (1 if error_mode == 'output' else 3 if error_mode == 'backward' else 4)
|
||||
|
||||
for name, data in results.items():
|
||||
avg_diff = data['sum_diff'] / total_elements
|
||||
max_diff = data['max_diff']
|
||||
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
|
||||
|
||||
return results
|
||||
|
||||
def generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward'):
|
||||
seq_lengths = [768 * (2**i) for i in range(1)]
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print(f"MHA ERROR COMPARISON TABLE (b={b}, h={h}, d={d}, mean={mean}, std={std})")
|
||||
print(f"Mode: {error_mode}, Test: {test_mode}")
|
||||
print(f"{'='*80}")
|
||||
|
||||
# Print header
|
||||
print(f"{'Seq Length':<12} | {'MHA vs PT Avg':<15} | {'MHA vs PT Max':<15} | {'FA2 vs PT Avg':<15} | {'FA2 vs PT Max':<15}")
|
||||
print(f"{'-'*12} | {'-'*15} | {'-'*15} | {'-'*15} | {'-'*15}")
|
||||
|
||||
for n in seq_lengths:
|
||||
results = check_correctness(b, h, n, d, mean, std, error_mode=error_mode, test_mode=test_mode)
|
||||
|
||||
mha_pt_avg = results['MHA vs PT']['avg_diff']
|
||||
mha_pt_max = results['MHA vs PT']['max_diff']
|
||||
fa2_pt_avg = results['FA2 vs PT']['avg_diff']
|
||||
fa2_pt_max = results['FA2 vs PT']['max_diff']
|
||||
|
||||
# Print row
|
||||
print(f"{n:<12} | {mha_pt_avg:<15.6e} | {mha_pt_max:<15.6e} | {fa2_pt_avg:<15.6e} | {fa2_pt_max:<15.6e}")
|
||||
|
||||
print(f"{'='*80}\n")
|
||||
|
||||
# fix random seed
|
||||
torch.manual_seed(0)
|
||||
|
||||
# Example usage
|
||||
b, h, d = 2, 2, 64
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
# Test forward only
|
||||
generate_error_tables(b, h, d, mean, std, error_mode='output', test_mode='forward_only')
|
||||
|
||||
# Test forward and backward
|
||||
generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward')
|
||||
|
||||
print("MHA attention error comparison completed.")
|
||||
@@ -1,27 +0,0 @@
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/ATen.h>
|
||||
|
||||
#include <vector>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num
|
||||
);
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_backward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num
|
||||
);
|
||||
#endif
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
|
||||
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
|
||||
#endif
|
||||
}
|
||||
@@ -1,470 +0,0 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch.utils.checkpoint import detach_variable
|
||||
from typing import Tuple
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
|
||||
def video_sparse_attn(q, k, v, topk, block_size, compress_attn_weight=None):
|
||||
"""
|
||||
q: [batch_size, num_heads, seq_len, head_dim]
|
||||
k: [batch_size, num_heads, seq_len, head_dim]
|
||||
v: [batch_size, num_heads, seq_len, head_dim]
|
||||
topk: int
|
||||
block_size: int or tuple of 3 ints
|
||||
video_shape: tuple of (T, H, W)
|
||||
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
|
||||
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
|
||||
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
|
||||
"""
|
||||
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
assert block_elements % 64 == 0 and block_elements >= 64
|
||||
assert q.shape[2] % block_elements == 0
|
||||
batch_size, num_heads, seq_len, head_dim = q.shape
|
||||
# compress attn
|
||||
q_compress = q.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
k_compress = k.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
v_compress = v.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
|
||||
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
|
||||
v_compress)
|
||||
|
||||
output_compress = output_compress.view(batch_size, num_heads,
|
||||
seq_len // block_elements, 1,
|
||||
head_dim)
|
||||
output_compress = output_compress.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch_size, num_heads,
|
||||
seq_len, head_dim)
|
||||
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num = generate_topk_block_sparse_pattern(
|
||||
block_attn_score, topk)
|
||||
|
||||
output_select = block_sparse_attn(q, k, v, q2k_block_sparse_index,
|
||||
q2k_block_sparse_num,
|
||||
k2q_block_sparse_index,
|
||||
k2q_block_sparse_num)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
final_output = output_compress * compress_attn_weight + output_select
|
||||
else:
|
||||
final_output = output_compress + output_select
|
||||
return final_output
|
||||
|
||||
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
QK = torch.matmul(q, k.transpose(-2, -1))
|
||||
QK /= (q.size(-1)**0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v)
|
||||
return output, QK
|
||||
|
||||
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
|
||||
topk: int):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
|
||||
based on the provided attention scores.
|
||||
|
||||
Args:
|
||||
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Attention scores between query and key blocks
|
||||
topk: int
|
||||
Number of kv blocks each q block attends to
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to topk).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
"""
|
||||
device = block_attn_score.device
|
||||
# Extract dimensions from block_attn_score
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
|
||||
|
||||
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
|
||||
|
||||
sorted_indice = sorted_result.indices
|
||||
|
||||
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
|
||||
dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
|
||||
topk,
|
||||
device=device,
|
||||
dtype=torch.int32)
|
||||
|
||||
block_map = topk_index_to_map(q2k_block_sparse_index,
|
||||
num_kv_blocks,
|
||||
transpose_map=True)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
|
||||
block_map.transpose(2, 3))
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
|
||||
@torch._dynamo.disable
|
||||
def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
"""
|
||||
Differentiable block sparse attention function.
|
||||
|
||||
Args:
|
||||
q: Query tensor [batch_size, num_heads, seq_len_q, head_dim]
|
||||
k: Key tensor [batch_size, num_heads, seq_len_kv, head_dim]
|
||||
v: Value tensor [batch_size, num_heads, seq_len_kv, head_dim]
|
||||
q2k_block_sparse_index: Indices for query-to-key sparse blocks
|
||||
q2k_block_sparse_num: Number of sparse blocks for each query block
|
||||
k2q_block_sparse_index: Indices for key-to-query sparse blocks (for backward pass)
|
||||
k2q_block_sparse_num: Number of sparse blocks for each key block (for backward pass)
|
||||
|
||||
Returns:
|
||||
output: Attention output tensor [batch_size, num_heads, seq_len_q, head_dim]
|
||||
"""
|
||||
return BlockSparseAttentionFunction.apply(
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
)
|
||||
|
||||
def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num):
|
||||
"""
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks].
|
||||
[*, *, i, j] = 1 means the i-th q block should attend to the j-th kv block.
|
||||
"""
|
||||
# assert all elements in q2k_block_sparse_num can be devisible by 2
|
||||
o, lse = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
return o, lse
|
||||
|
||||
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
grad_output = grad_output.contiguous()
|
||||
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return grad_q, grad_k, grad_v
|
||||
|
||||
## pytorch sdpa version of block sparse ##
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
@triton.jit
|
||||
def index_to_mask_kernel(
|
||||
q2k_block_sparse_index_ptr,
|
||||
q2k_block_sparse_num_ptr,
|
||||
mask_ptr,
|
||||
batch_size: tl.constexpr,
|
||||
num_heads: tl.constexpr,
|
||||
num_q_blocks: tl.constexpr,
|
||||
num_k_blocks: tl.constexpr,
|
||||
max_kv_blocks: tl.constexpr,
|
||||
BLOCK_Q: tl.constexpr,
|
||||
BLOCK_K: tl.constexpr,
|
||||
):
|
||||
bh, q, id = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
|
||||
b = bh // num_heads
|
||||
h = bh % num_heads
|
||||
|
||||
num_valid_blocks = tl.load(q2k_block_sparse_num_ptr + b * num_heads * num_q_blocks + h * num_q_blocks + q)
|
||||
|
||||
if num_valid_blocks <= id:
|
||||
return
|
||||
k = tl.load(q2k_block_sparse_index_ptr + b * num_heads * num_q_blocks * max_kv_blocks + h * num_q_blocks * max_kv_blocks + q * max_kv_blocks + id)
|
||||
|
||||
full_mask = (tl.arange(0, BLOCK_Q)[:, None] < BLOCK_Q) & (tl.arange(0, BLOCK_K)[None, :] < BLOCK_K)
|
||||
|
||||
q_lengths = num_q_blocks * BLOCK_Q
|
||||
k_lengths = num_k_blocks * BLOCK_K
|
||||
mask_ptr_base = mask_ptr + b * num_heads * q_lengths * k_lengths + h * q_lengths * k_lengths + q * BLOCK_Q * k_lengths + k * BLOCK_K
|
||||
|
||||
tl.store(mask_ptr_base + tl.arange(0, BLOCK_Q)[:, None] * k_lengths + tl.arange(0, BLOCK_K)[None, :], full_mask)
|
||||
|
||||
def index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, BLOCK_Q, BLOCK_K, num_k_blocks):
|
||||
"""
|
||||
Convert block sparse indices to a mask.
|
||||
|
||||
Args:
|
||||
q2k_block_sparse_index: Indices for query-to-key sparse blocks
|
||||
q2k_block_sparse_num: Number of sparse blocks for each query block
|
||||
|
||||
Returns:
|
||||
mask: Block sparse mask tensor
|
||||
"""
|
||||
batch_size, num_heads, num_q_blocks, max_kv_blocks = q2k_block_sparse_index.shape
|
||||
assert q2k_block_sparse_num.shape == (batch_size, num_heads, num_q_blocks)
|
||||
|
||||
mask = torch.zeros((batch_size, num_heads, num_q_blocks * BLOCK_Q, num_k_blocks * BLOCK_K), dtype=torch.bool, device=q2k_block_sparse_index.device)
|
||||
|
||||
grid = (batch_size * num_heads, num_q_blocks, max_kv_blocks)
|
||||
index_to_mask_kernel[grid](
|
||||
q2k_block_sparse_index,
|
||||
q2k_block_sparse_num,
|
||||
mask,
|
||||
batch_size,
|
||||
num_heads,
|
||||
num_q_blocks,
|
||||
num_k_blocks,
|
||||
max_kv_blocks,
|
||||
BLOCK_Q=BLOCK_Q,
|
||||
BLOCK_K=BLOCK_K,
|
||||
)
|
||||
|
||||
return mask
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
topk: tl.constexpr,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
for i in tl.static_range(topk):
|
||||
index = tl.load(index_ptr_base + i * index_kv_stride)
|
||||
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
|
||||
|
||||
@triton.jit
|
||||
def map_to_index_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
index_num_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
index_num_bs_stride,
|
||||
index_num_h_stride,
|
||||
index_num_q_stride,
|
||||
num_kv_blocks: tl.constexpr,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
num = 0
|
||||
for i in tl.static_range(num_kv_blocks):
|
||||
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
|
||||
if map_entry:
|
||||
tl.store(index_ptr_base + num * index_kv_stride, i)
|
||||
num += 1
|
||||
|
||||
tl.store(
|
||||
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
|
||||
q * index_num_q_stride, num)
|
||||
|
||||
def topk_index_to_map(index: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
transpose_map: bool = False):
|
||||
"""
|
||||
Convert topk indices to a map.
|
||||
|
||||
Args:
|
||||
index: [bs, h, num_q_blocks, topk]
|
||||
The topk indices tensor.
|
||||
num_kv_blocks: int
|
||||
The number of key-value blocks in the block_map returned
|
||||
transpose_map: bool
|
||||
If True, the block_map will be transposed on the final two dimensions.
|
||||
|
||||
Returns:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
A binary map where 1 indicates that the q block attends to the kv block.
|
||||
"""
|
||||
bs, h, num_q_blocks, topk = index.shape
|
||||
|
||||
if transpose_map is False:
|
||||
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
else:
|
||||
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
block_map = block_map.transpose(2, 3)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
topk_index_to_map_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
topk=topk,
|
||||
)
|
||||
|
||||
return block_map
|
||||
|
||||
def map_to_index(block_map: torch.Tensor):
|
||||
"""
|
||||
Convert a block map to indices and counts.
|
||||
|
||||
Args:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The block map tensor.
|
||||
|
||||
Returns:
|
||||
index: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The indices of the blocks.
|
||||
index_num: [bs, h, num_q_blocks]
|
||||
The number of blocks for each q block.
|
||||
"""
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
|
||||
|
||||
index = torch.full((block_map.shape),
|
||||
-1,
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
index_num = torch.empty((bs, h, num_q_blocks),
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
map_to_index_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
index_num,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
index_num.stride(0),
|
||||
index_num.stride(1),
|
||||
index_num.stride(2),
|
||||
num_kv_blocks=num_kv_blocks,
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
|
||||
class BlockSparseAttentionFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
o, lse = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
ctx.save_for_backward(q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num = ctx.saved_tensors
|
||||
grad_q, grad_k, grad_v = block_sparse_attention_backward(
|
||||
q, k, v, o, lse, grad_output, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
)
|
||||
return grad_q, grad_k, grad_v, None, None, None, None
|
||||
|
||||
|
||||
class DummyOperator(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x):
|
||||
return x
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
return grad_output
|
||||
|
||||
class CheckpointSDPA(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, obj, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k):
|
||||
"""Forward pass."""
|
||||
with torch.no_grad():
|
||||
mask = index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k, k.shape[2] // block_k)
|
||||
outputs = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
|
||||
ctx.save_for_backward(*detach_variable((q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)))
|
||||
ctx.block_q = block_q
|
||||
ctx.block_k = block_k
|
||||
# the obj is passed in, then it can access the saved input
|
||||
# tensors later for recomputation
|
||||
obj.ctx = ctx
|
||||
return outputs
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
"""Backward pass."""
|
||||
inputs = ctx.saved_tensors
|
||||
output = ctx.output
|
||||
torch.autograd.backward(output, grad_output)
|
||||
ctx.output = None
|
||||
grads = tuple(inp.grad for inp in inputs)
|
||||
return (None, ) + grads + (None, None)
|
||||
|
||||
|
||||
class BlockSparseAttnTorch:
|
||||
def __init__(self):
|
||||
self.ctx = None
|
||||
|
||||
def recompute_mask(self, _):
|
||||
recomputed_mask = index_to_mask(self.q2k_block_sparse_index, self.q2k_block_sparse_num, self.block_q, self.block_k, self.num_kv_blocks)
|
||||
mask_size = recomputed_mask.untyped_storage().size()
|
||||
self.mask.untyped_storage().resize_(mask_size)
|
||||
self.mask.untyped_storage().copy_(recomputed_mask.untyped_storage())
|
||||
|
||||
def recompute(self, _):
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num = self.ctx.saved_tensors
|
||||
block_q = self.ctx.block_q
|
||||
block_k = self.ctx.block_k
|
||||
mask = index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k, k.shape[2] // block_k)
|
||||
with torch.enable_grad():
|
||||
output = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
|
||||
self.ctx.output = output
|
||||
self.ctx = None
|
||||
|
||||
@torch._dynamo.disable
|
||||
def forward(self, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k):
|
||||
"""
|
||||
Differentiable block sparse attention function using PyTorch.
|
||||
|
||||
Args:
|
||||
q: Query tensor [batch_size, num_heads, seq_len_q, head_dim]
|
||||
k: Key tensor [batch_size, num_heads, seq_len_kv, head_dim]
|
||||
v: Value tensor [batch_size, num_heads, seq_len_kv, head_dim]
|
||||
q2k_block_sparse_index: Indices for query-to-key sparse blocks
|
||||
q2k_block_sparse_num: Number of sparse blocks for each query block
|
||||
block_q: Block size for query
|
||||
block_k: Block size for key-value
|
||||
|
||||
Returns:
|
||||
output: Attention output tensor [batch_size, num_heads, seq_len_q, head_dim]
|
||||
"""
|
||||
|
||||
output = CheckpointSDPA.apply(
|
||||
self, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k
|
||||
)
|
||||
|
||||
o = DummyOperator.apply(output)
|
||||
o.register_hook(self.recompute)
|
||||
return o
|
||||
File diff suppressed because it is too large
Load Diff
@@ -6,6 +6,7 @@
|
||||
## Installation
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
|
||||
First, install C++20 for ThunderKittens:
|
||||
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
@@ -15,27 +16,17 @@ sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
|
||||
## Environment Setup
|
||||
First, set up your CUDA environment:
|
||||
Install STA:
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
```
|
||||
|
||||
## Install Sliding Tile Attention (STA)
|
||||
```bash
|
||||
python setup_sta.py install
|
||||
```
|
||||
|
||||
## Install Video Sparse Attention (VSA)
|
||||
```bash
|
||||
python setup_vsa.py install
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
@@ -1,6 +1,6 @@
|
||||
### ADD TO THIS TO REGISTER NEW KERNELS
|
||||
sources = {
|
||||
'st_attn': {
|
||||
'attn': {
|
||||
'source_files': {
|
||||
'h100': 'st_attn/st_attn_h100.cu' # define these source files for each GPU target desired.
|
||||
}
|
||||
@@ -9,7 +9,7 @@ sources = {
|
||||
|
||||
### WHICH KERNELS DO WE WANT TO BUILD?
|
||||
# (oftentimes during development work you don't need to redefine them all.)
|
||||
kernels = ['st_attn']
|
||||
kernels = ['attn']
|
||||
|
||||
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
|
||||
target = 'h100'
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from csrc.attn.config_sta import kernels, sources, target
|
||||
from config import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
@@ -7,7 +7,8 @@
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
|
||||
#ifdef TK_COMPILE_ATTN
|
||||
extern torch::Tensor sta_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
|
||||
);
|
||||
@@ -16,8 +17,8 @@ extern torch::Tensor sta_forward(
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
|
||||
#ifdef TK_COMPILE_ATTN
|
||||
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
|
||||
#endif
|
||||
}
|
||||
}
|
||||
@@ -1,22 +1,19 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch.utils.checkpoint import detach_variable
|
||||
try:
|
||||
from st_attn_cuda import sta_fwd
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
from st_attn_cuda import sta_fwd
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, img_latent_shape='30*48*80'):
|
||||
seq_length = q_all.shape[2]
|
||||
dit_seq_shape_mapping = {
|
||||
img_latent_shape_mapping = {
|
||||
'30x48x80':1,
|
||||
'36x48x48':2,
|
||||
'18x48x80':3,
|
||||
}
|
||||
if has_text:
|
||||
assert q_all.shape[
|
||||
2] >= 115200 and q_all.shape[2] <= 115456, f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '30x48x80' for HunyuanVideo"
|
||||
2] >= 115200, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
|
||||
assert q_all.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
|
||||
target_size = math.ceil(seq_length / 384) * 384
|
||||
pad_size = target_size - seq_length
|
||||
@@ -25,14 +22,14 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
|
||||
k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
|
||||
v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
|
||||
else:
|
||||
if dit_seq_shape == '36x48x48': # Stepvideo 204x768x68
|
||||
if img_latent_shape == '36x48x48': # Stepvideo 204x768x68
|
||||
assert q_all.shape[2] == 82944
|
||||
elif dit_seq_shape == '18x48x80': # Wan 69x768x1280
|
||||
elif img_latent_shape == '18x48x80': # Wan 69x768x1280
|
||||
assert q_all.shape[2] == 69120
|
||||
else:
|
||||
raise ValueError(f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
|
||||
raise ValueError(f"Unsupported {img_latent_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
|
||||
|
||||
kernel_aspect_ratio_flag = dit_seq_shape_mapping[dit_seq_shape]
|
||||
kernel_aspect_ratio_flag = img_latent_shape_mapping[img_latent_shape]
|
||||
hidden_states = torch.empty_like(q_all)
|
||||
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
|
||||
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
|
||||
@@ -46,4 +43,4 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
|
||||
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
|
||||
if has_text:
|
||||
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
|
||||
return hidden_states[:, :, :seq_length]
|
||||
return hidden_states[:, :, :seq_length]
|
||||
-1
@@ -829,4 +829,3 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
return o;
|
||||
cudaDeviceSynchronize();
|
||||
}
|
||||
|
||||
@@ -45,12 +45,12 @@ def benchmark_attention(configurations):
|
||||
|
||||
# Warmup for forward pass
|
||||
for _ in range(10):
|
||||
o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, '18x48x80')
|
||||
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
|
||||
|
||||
# Time the forward pass
|
||||
for i in range(10):
|
||||
start_events_fwd[i].record()
|
||||
o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, '18x48x80')
|
||||
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
|
||||
end_events_fwd[i].record()
|
||||
|
||||
torch.cuda.synchronize()
|
||||
@@ -124,7 +124,7 @@ def plot_results(results):
|
||||
|
||||
# Example list of configurations to test
|
||||
configurations = [
|
||||
(2, 24, 69120, 128, False),
|
||||
(2, 24, 82944, 128, False),
|
||||
# (16, 16, 768*16, 128, False),
|
||||
# (16, 16, 768*2, 128, False),
|
||||
# (16, 16, 768*4, 128, False),
|
||||
@@ -9,14 +9,14 @@ flex_attention = torch.compile(flex_attention, dynamic=False)
|
||||
|
||||
|
||||
def flex_test(Q, K, V, kernel_size):
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (36, 48, 48), 39, 'cuda', 0)
|
||||
output = flex_attention(Q, K, V, block_mask=mask)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def h100_fwd_kernel_test(Q, K, V, kernel_size):
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 39, False)
|
||||
return o
|
||||
|
||||
|
||||
@@ -37,7 +37,7 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
|
||||
'max_diff': 0
|
||||
},
|
||||
}
|
||||
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
|
||||
kernel_size_ls = [(6, 1, 6), (6, 6, 1)]
|
||||
from tqdm import tqdm
|
||||
for kernel_size in tqdm(kernel_size_ls):
|
||||
for _ in range(num_iterations):
|
||||
@@ -72,14 +72,25 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
|
||||
return results
|
||||
|
||||
|
||||
def generate_error_graphs(b, h, d, causal, mean, std, error_mode='all'):
|
||||
seq_lengths = [82944]
|
||||
|
||||
tk_avg_errors, tk_max_errors = [], []
|
||||
|
||||
for n in tqdm(seq_lengths, desc="Generating error data"):
|
||||
results = check_correctness(b, h, n, d, causal, mean, std, error_mode=error_mode)
|
||||
|
||||
tk_avg_errors.append(results['TK vs FLEX']['avg_diff'])
|
||||
tk_max_errors.append(results['TK vs FLEX']['max_diff'])
|
||||
|
||||
|
||||
# Example usage
|
||||
b, h, d = 2, 24, 128
|
||||
n = 69120 # Sequence length
|
||||
causal = False
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
# Run correctness check directly
|
||||
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
|
||||
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
|
||||
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
|
||||
for mode in ['output']:
|
||||
generate_error_graphs(b, h, d, causal, mean, std, error_mode=mode)
|
||||
|
||||
print("Error graphs generated and saved for all modes.")
|
||||
@@ -0,0 +1,46 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
|
||||
DATA_DIR=./data
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
# --gradient_checkpointing\
|
||||
# --pretrained_model_name_or_path hunyuanvideo-community/HunyuanVideo \
|
||||
# --pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
torchrun --nnodes 1 --nproc_per_node 4\
|
||||
fastvideo/v1/pipelines/training_pipeline.py\
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache"\
|
||||
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 1 \
|
||||
--sp_size 4 \
|
||||
--tp_size 4 \
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=320\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_16_HD"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_height 720 \
|
||||
--num_width 1280 \
|
||||
--num_frames 125 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver \
|
||||
--master_weight_type "bf16"
|
||||
@@ -57,9 +57,8 @@ Run the script with:
|
||||
python example.py
|
||||
```
|
||||
|
||||
The generated video will be saved in the current directory under `my_videos/`
|
||||
The generated video will be saved in the current directory under `my_videos/`.
|
||||
|
||||
More inference example scripts can be found in `scripts/inference/`
|
||||
## Available Models
|
||||
|
||||
Please see the [support matrix](#support-matrix) for the list of supported models and their available optimizations.
|
||||
@@ -80,6 +79,7 @@ def main():
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param.num_frames = 107
|
||||
sampling_param.image_strength = 0.8 # How much to preserve the original image (0-1)
|
||||
|
||||
# Generate video based on the image
|
||||
prompt = "A photograph coming to life with gentle movement"
|
||||
|
||||
@@ -7,40 +7,70 @@ To save GPU memory, we precompute text embeddings and VAE latents to eliminate t
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=FastVideo/mini_i2v_dataset --repo_type=dataset
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
|
||||
```
|
||||
|
||||
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
|
||||
|
||||
To preprocess the dataset for fine-tuning or distillation, run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/v1_preprocess_wan_data_t2v # for wan
|
||||
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
|
||||
bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
|
||||
```
|
||||
|
||||
The preprocessed dataset will be stored in `Image-Vid-Finetune-Mochi` or `Image-Vid-Finetune-HunYuan` correspondingly.
|
||||
|
||||
## Process your own dataset
|
||||
|
||||
If you wish to create your own dataset for finetuning or distillation, please refer `mini_i2v_dataset/crush-smol_raw/` to structure you video dataset in the following format:
|
||||
If you wish to create your own dataset for finetuning or distillation, please structure you video dataset in the following format:
|
||||
|
||||
```
|
||||
path_to_your_dataset_folder/
|
||||
├── videos/
|
||||
│ ├── 0.mp4
|
||||
path_to_dataset_folder/
|
||||
├── media/
|
||||
│ ├── 0.jpg
|
||||
│ ├── 1.mp4
|
||||
├── videos.txt
|
||||
└── prompt.txt
|
||||
│ ├── 2.jpg
|
||||
├── video2caption.json
|
||||
└── merge.txt
|
||||
```
|
||||
|
||||
To geranate the `videos2caption.json` and `merge.txt`, run
|
||||
Format the JSON file as a list, where each item represents a media source:
|
||||
|
||||
``` python
|
||||
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
|
||||
```
|
||||
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/v1_preprocess_****.sh` accordingly and run:
|
||||
For image media,
|
||||
|
||||
```
|
||||
bash scripts/preprocess/v1_preprocess_****.sh
|
||||
{
|
||||
"path": "0.jpg",
|
||||
"cap": ["captions"]
|
||||
}
|
||||
```
|
||||
|
||||
For video media,
|
||||
|
||||
```
|
||||
{
|
||||
"path": "1.mp4",
|
||||
"resolution": {
|
||||
"width": 848,
|
||||
"height": 480
|
||||
},
|
||||
"fps": 30.0,
|
||||
"duration": 6.033333333333333,
|
||||
"cap": [
|
||||
"caption"
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
Use a txt file (merge.txt) to contain the source folder for media and the JSON file for meta information:
|
||||
|
||||
```
|
||||
path_to_media_source_foder,path_to_json_file
|
||||
```
|
||||
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/preprocess_****_data.sh` accordingly and run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/preprocess_****_data.sh
|
||||
```
|
||||
|
||||
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
|
||||
|
||||
@@ -2,7 +2,7 @@ from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -11,9 +11,7 @@ def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# if num_gpus > 1, FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
use_cpu_offload=False
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
@@ -25,7 +23,7 @@ def main():
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
video = generator.generate_video(prompt)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
@@ -36,7 +34,7 @@ def main():
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
video2 = generator.generate_video(prompt2)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
|
||||
@@ -1,45 +0,0 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora"
|
||||
def main():
|
||||
# Initialize VideoGenerator with the Wan model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=2,
|
||||
lora_path="benjamin-paine/steamboat-willie-1.3b",
|
||||
lora_nickname="steamboat"
|
||||
)
|
||||
kwargs = {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 32,
|
||||
}
|
||||
# Generate video with LoRA style
|
||||
prompt = "steamboat willie style, golden era animation, close-up of a short fluffy monster kneeling beside a melting red candle. the mood is one of wonder and curiosity, as the monster gazes at the flame with wide eyes and open mouth. Its pose and expression convey a sense of innocence and playfulness, as if it is exploring the world around it for the first time. The use of warm colors and dramatic lighting further enhances the cozy atmosphere of the image."
|
||||
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
# sampling_param=sampling_param,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
generator.set_lora_adapter(lora_nickname="flat_color", lora_path="motimalu/wan-flat-color-1.3b-v2")
|
||||
prompt = "flat color, no lineart, blending, negative space, artist:[john kafka|ponsuke kaikai|hara id 21|yoneyama mai|fuzichoco], 1girl, sakura miko, pink hair, cowboy shot, white shirt, floral print, off shoulder, outdoors, cherry blossom, tree shade, wariza, looking up, falling petals, half-closed eyes, white sky, clouds, live2d animation, upper body, high quality cinematic video of a woman sitting under a sakura tree. Dreamy and lonely, the camera close-ups on the face of the woman as she turns towards the viewer. The Camera is steady, This is a cowboy shot. The animation is smooth and fluid."
|
||||
negative_prompt = "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,5 +0,0 @@
|
||||
# STA Mask Search Examples
|
||||
|
||||
```bash
|
||||
bash examples/inference/sta_mask_search/inference_wan_sta.sh
|
||||
```
|
||||
@@ -1,39 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
|
||||
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
|
||||
base_port=29503
|
||||
num_gpu=$(nvidia-smi --query-gpu=gpu_name --format=csv,noheader | wc -l)
|
||||
gpu_ids=$(seq 0 $((num_gpu-1)))
|
||||
skip_time_steps=12
|
||||
|
||||
output_path="inference_results/sta/mask_search_full"
|
||||
STA_mode="STA_searching"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_extend_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode &
|
||||
sleep 1
|
||||
done
|
||||
wait
|
||||
echo "STA searching completed"
|
||||
|
||||
output_path="inference_results/sta/mask_search_sparse"
|
||||
STA_mode="STA_tuning"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_extend_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode \
|
||||
--skip_time_steps $skip_time_steps &
|
||||
sleep 1
|
||||
done
|
||||
wait
|
||||
echo "STA tuning completed"
|
||||
|
||||
echo "All jobs completed"
|
||||
@@ -1,63 +0,0 @@
|
||||
import os
|
||||
import argparse
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
def main(args):
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
num_gpus=args.num_gpus, # Adjust based on your hardware
|
||||
STA_mode=args.STA_mode,
|
||||
skip_time_steps=args.skip_time_steps
|
||||
)
|
||||
|
||||
# Prompts for your video
|
||||
prompt = args.prompt
|
||||
prompt_path = args.prompt_path
|
||||
negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
|
||||
if prompt_path is not None:
|
||||
with open(prompt_path, "r") as f:
|
||||
prompts = f.readlines()
|
||||
else:
|
||||
prompts = [prompt]
|
||||
|
||||
params = SamplingParam(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
fps=args.fps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
return_frames=True, # Also return frames from this call (defaults to False)
|
||||
output_path=args.output_path, # Controls where videos are saved
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt
|
||||
)
|
||||
|
||||
# Generate the video
|
||||
for prompt in prompts:
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
sampling_param=params,
|
||||
)
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompt", type=str, default="A man is dancing.")
|
||||
parser.add_argument("--prompt_path", type=str, default=None)
|
||||
parser.add_argument("--height", type=int, default=768)
|
||||
parser.add_argument("--width", type=int, default=1280)
|
||||
parser.add_argument("--num_frames", type=int, default=69)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=50)
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--guidance_scale", type=float, default=5.0)
|
||||
parser.add_argument("--seed", type=int, default=12345)
|
||||
parser.add_argument("--output_path", type=str, default="my_videos/")
|
||||
parser.add_argument("--num_gpus", type=int, default=1)
|
||||
parser.add_argument("--STA_mode", type=str, default="STA_searching")
|
||||
parser.add_argument("--skip_time_steps", type=int, default=12)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -1,6 +1,5 @@
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.version import __version__
|
||||
|
||||
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
|
||||
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam"]
|
||||
|
||||
@@ -0,0 +1,122 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import maybe_download_model, shallow_asdict
|
||||
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo.v1.pipelines.preprocess_pipeline import PreprocessPipeline
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
BASE_MODEL_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
|
||||
def main(args):
|
||||
# Assume using torchrun
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
init_distributed_environment(world_size=world_size, rank=rank, local_rank=local_rank)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
|
||||
torch.cuda.set_device(local_rank)
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(MODEL_PATH)
|
||||
kwargs = {
|
||||
"use_cpu_offload": False,
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
|
||||
}
|
||||
pipeline_config_args = shallow_asdict(pipeline_config)
|
||||
pipeline_config_args.update(kwargs)
|
||||
fastvideo_args = FastVideoArgs(model_path=MODEL_PATH,
|
||||
num_gpus=world_size,
|
||||
device_str="cuda",
|
||||
**pipeline_config_args,
|
||||
)
|
||||
fastvideo_args.check_fastvideo_args()
|
||||
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
|
||||
|
||||
pipeline = PreprocessPipeline(MODEL_PATH, fastvideo_args)
|
||||
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--model_type", type=str, default="mochi")
|
||||
parser.add_argument("--data_merge_path", type=str, required=True)
|
||||
parser.add_argument("--validation_prompt_txt", type=str)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument(
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--preprocess_video_batch_size",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--preprocess_text_batch_size",
|
||||
type=int,
|
||||
default=8,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--samples_per_file",
|
||||
type=int,
|
||||
default=64
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flush_frequency",
|
||||
type=int,
|
||||
default=256,
|
||||
help="how often to save to parquet files"
|
||||
)
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument("--max_height", type=int, default=480)
|
||||
parser.add_argument("--max_width", type=int, default=848)
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--dataset", default="t2v")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
parser.add_argument("--speed_factor", type=float, default=1.0)
|
||||
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
parser.add_argument("--cfg", type=float, default=0.0)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -68,8 +68,7 @@ def main(args):
|
||||
train_dataset = T5dataset(latents_json_path, args.vae_debug)
|
||||
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
vae.enable_tiling()
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
|
||||
@@ -0,0 +1,199 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from accelerate.logging import get_logger
|
||||
from diffusers.utils import export_to_video
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm import tqdm
|
||||
|
||||
# from fastvideo.utils.load import load_text_encoder, load_vae
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import VAELoader, TextEncoderLoader, TokenizerLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.configs.models.encoders.t5 import T5Config
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
class T5dataset(Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
json_path,
|
||||
vae_debug,
|
||||
):
|
||||
self.json_path = json_path
|
||||
self.vae_debug = vae_debug
|
||||
with open(self.json_path, "r") as f:
|
||||
train_dataset = json.load(f)
|
||||
self.train_dataset = sorted(train_dataset, key=lambda x: x["latent_path"])
|
||||
|
||||
def __getitem__(self, idx):
|
||||
caption = self.train_dataset[idx]["caption"]
|
||||
filename = self.train_dataset[idx]["latent_path"].split(".")[0]
|
||||
length = self.train_dataset[idx]["length"]
|
||||
if self.vae_debug:
|
||||
latents = torch.load(
|
||||
os.path.join(args.output_dir, "latent", self.train_dataset[idx]["latent_path"]),
|
||||
map_location="cpu",
|
||||
)
|
||||
else:
|
||||
latents = []
|
||||
|
||||
return dict(caption=caption, latents=latents, filename=filename, length=length)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.train_dataset)
|
||||
|
||||
|
||||
def main(args):
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
rank = int(os.getenv("RANK", 0))
|
||||
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
|
||||
print("world_size", world_size, "local rank", local_rank)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
torch.cuda.set_device(device)
|
||||
world_group = get_world_group()
|
||||
|
||||
# device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
# torch.cuda.set_device(local_rank)
|
||||
# if not dist.is_initialized():
|
||||
# dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
|
||||
videoprocessor = VideoProcessor(vae_scale_factor=8)
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "video"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "prompt_embed"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "prompt_attention_mask"), exist_ok=True)
|
||||
|
||||
vae_precision = "fp16"
|
||||
text_encoder_precision = "fp32"
|
||||
fastvideo_args = FastVideoArgs(model_path=args.model_path,
|
||||
use_cpu_offload=False,
|
||||
vae_precision=vae_precision,
|
||||
text_encoder_precisions=(text_encoder_precision,))
|
||||
fastvideo_args.device = device
|
||||
fastvideo_args.device_str = f"cuda:{local_rank}"
|
||||
|
||||
# fastvideo_args.dit_config = HunyuanVideoConfig()
|
||||
fastvideo_args.vae_config = WanVAEConfig()
|
||||
fastvideo_args.text_encoder_configs = (T5Config(),)
|
||||
|
||||
# vae_loader = VAELoader()
|
||||
# vae = vae_loader.load_vae()
|
||||
text_encoder_loader = TextEncoderLoader()
|
||||
tokenizer_loader = TokenizerLoader()
|
||||
|
||||
model_path = args.model_path
|
||||
path = maybe_download_model(model_path)
|
||||
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
ENCODER_PATH = os.path.join(path, "text_encoder")
|
||||
TOKENIZER_PATH = os.path.join(path, "tokenizer")
|
||||
print(ENCODER_PATH)
|
||||
text_encoder = text_encoder_loader.load(ENCODER_PATH, "text_encoder", fastvideo_args)
|
||||
tokenizer = tokenizer_loader.load(TOKENIZER_PATH, "tokenizer", fastvideo_args)
|
||||
|
||||
latents_json_path = os.path.join(args.output_dir, "videos2caption_temp.json")
|
||||
train_dataset = T5dataset(latents_json_path, args.vae_debug)
|
||||
# text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
|
||||
# vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
# vae.enable_tiling()
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
json_data = []
|
||||
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
|
||||
with torch.inference_mode():
|
||||
# with torch.autocast("cuda", dtype=torch.float32):
|
||||
print(data["caption"])
|
||||
text_inputs = tokenizer(data["caption"], **fastvideo_args.text_encoder_configs[0].tokenizer_kwargs).to(
|
||||
fastvideo_args.device)
|
||||
input_ids = text_inputs["input_ids"]
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
outputs = text_encoder(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
from fastvideo.v1.configs.pipelines.wan import t5_postprocess_text
|
||||
post_process_func = t5_postprocess_text
|
||||
prompt_embeds = post_process_func(outputs)
|
||||
prompt_attention_mask = attention_mask
|
||||
if args.vae_debug:
|
||||
latents = data["latents"]
|
||||
video = vae.decode(latents.to(device), return_dict=False)[0]
|
||||
video = videoprocessor.postprocess_video(video)
|
||||
for idx, video_name in enumerate(data["filename"]):
|
||||
prompt_embed_path = os.path.join(args.output_dir, "prompt_embed", video_name + ".pt")
|
||||
video_path = os.path.join(args.output_dir, "video", video_name + ".mp4")
|
||||
prompt_attention_mask_path = os.path.join(args.output_dir, "prompt_attention_mask",
|
||||
video_name + ".pt")
|
||||
# save latent
|
||||
torch.save(prompt_embeds[idx], prompt_embed_path)
|
||||
torch.save(prompt_attention_mask[idx], prompt_attention_mask_path)
|
||||
print(f"sample {video_name} saved")
|
||||
if args.vae_debug:
|
||||
export_to_video(video[idx], video_path, fps=16)
|
||||
item = {}
|
||||
item["length"] = int(data["length"][idx])
|
||||
item["latent_path"] = video_name + ".pt"
|
||||
item["prompt_embed_path"] = video_name + ".pt"
|
||||
item["prompt_attention_mask"] = video_name + ".pt"
|
||||
item["caption"] = data["caption"][idx]
|
||||
json_data.append(item)
|
||||
dist.barrier()
|
||||
local_data = json_data
|
||||
gathered_data = [None] * world_size
|
||||
dist.all_gather_object(gathered_data, local_data)
|
||||
if local_rank == 0:
|
||||
# os.remove(latents_json_path)
|
||||
all_json_data = [item for sublist in gathered_data for item in sublist]
|
||||
with open(os.path.join(args.output_dir, "videos2caption.json"), "w") as f:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
# parser.add_argument("--model_type", type=str, default="mochi")
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument(
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_batch_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument("--vae_debug", action="store_true")
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -33,8 +33,7 @@ def main(args):
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
vae.enable_tiling()
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
|
||||
@@ -104,7 +103,13 @@ if __name__ == "__main__":
|
||||
default=None,
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
# import torch.distributed as dist
|
||||
# from accelerate.logging import get_logger
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset import getdataset
|
||||
# from fastvideo.utils.load import load_vae
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import VAELoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
model_path = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
path = maybe_download_model(model_path)
|
||||
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
VAE_PATH = os.path.join(path, "vae")
|
||||
print(VAE_PATH)
|
||||
|
||||
|
||||
def main(args):
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
rank = int(os.getenv("RANK", 0))
|
||||
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
|
||||
print("world_size", world_size, "local rank", local_rank)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
torch.cuda.set_device(device)
|
||||
world_group = get_world_group()
|
||||
|
||||
vae_precision = "fp16"
|
||||
fastvideo_args = FastVideoArgs(model_path=VAE_PATH,
|
||||
use_cpu_offload=False,
|
||||
vae_precision=vae_precision)
|
||||
fastvideo_args.device = device
|
||||
# fastvideo_args.dit_config = HunyuanVideoConfig()
|
||||
fastvideo_args.vae_config = WanVAEConfig()
|
||||
|
||||
train_dataset = getdataset(args)
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
|
||||
# encoder_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
# torch.cuda.set_device(local_rank)
|
||||
# if not dist.is_initialized():
|
||||
# dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
vae_loader = VAELoader()
|
||||
vae = vae_loader.load(VAE_PATH, "vae", fastvideo_args)
|
||||
# vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
# vae.enable_tiling()
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
|
||||
json_data = []
|
||||
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
|
||||
with torch.inference_mode():
|
||||
with torch.autocast("cuda", dtype=torch.float16):
|
||||
latents = vae.encode(data["pixel_values"].to(device)).sample()
|
||||
for idx, video_path in enumerate(data["path"]):
|
||||
video_name = os.path.basename(video_path).split(".")[0]
|
||||
latent_path = os.path.join(args.output_dir, "latent", video_name + ".pt")
|
||||
torch.save(latents[idx].to(torch.bfloat16), latent_path)
|
||||
item = {}
|
||||
item["length"] = latents[idx].shape[1]
|
||||
item["latent_path"] = video_name + ".pt"
|
||||
item["caption"] = data["text"][idx]
|
||||
json_data.append(item)
|
||||
print(f"{video_name} processed")
|
||||
world_group.barrier()
|
||||
local_data = json_data
|
||||
gathered_data = [None] * world_size
|
||||
for i in range(world_size):
|
||||
if local_rank == i:
|
||||
world_group.broadcast_object(local_data, src=i)
|
||||
else:
|
||||
gathered_data[i] = world_group.broadcast_object(None, src=i)
|
||||
gathered_data[local_rank] = json_data
|
||||
print(gathered_data)
|
||||
if local_rank == 0:
|
||||
all_json_data = [item for sublist in gathered_data for item in sublist]
|
||||
with open(os.path.join(args.output_dir, "videos2caption_temp.json"), "w") as f:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
# parser.add_argument("--model_type", type=str, default="mochi")
|
||||
parser.add_argument("--data_merge_path", type=str, required=True)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument(
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_batch_size",
|
||||
type=int,
|
||||
default=16,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument("--max_height", type=int, default=480)
|
||||
parser.add_argument("--max_width", type=int, default=848)
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--dataset", default="t2v")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
parser.add_argument("--speed_factor", type=float, default=1.0)
|
||||
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
parser.add_argument("--cfg", type=float, default=0.0)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,115 @@
|
||||
import argparse
|
||||
import os
|
||||
|
||||
import torch
|
||||
# import torch.distributed as dist
|
||||
from accelerate.logging import get_logger
|
||||
|
||||
# from fastvideo.utils.load import load_text_encoder
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import VAELoader, TextEncoderLoader, TokenizerLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.configs.models.encoders.t5 import T5Config
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def main(args):
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
rank = int(os.getenv("RANK", 0))
|
||||
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
|
||||
print("world_size", world_size, "local rank", local_rank)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
torch.cuda.set_device(device)
|
||||
world_group = get_world_group()
|
||||
|
||||
vae_precision = "fp16"
|
||||
text_encoder_precision = "fp32"
|
||||
fastvideo_args = FastVideoArgs(model_path=args.model_path,
|
||||
use_cpu_offload=False,
|
||||
vae_precision=vae_precision,
|
||||
text_encoder_precisions=(text_encoder_precision,))
|
||||
fastvideo_args.device = device
|
||||
fastvideo_args.device_str = f"cuda:{local_rank}"
|
||||
|
||||
# fastvideo_args.dit_config = HunyuanVideoConfig()
|
||||
fastvideo_args.vae_config = WanVAEConfig()
|
||||
fastvideo_args.text_encoder_configs = (T5Config(),)
|
||||
|
||||
# vae_loader = VAELoader()
|
||||
# vae = vae_loader.load_vae()
|
||||
text_encoder_loader = TextEncoderLoader()
|
||||
tokenizer_loader = TokenizerLoader()
|
||||
|
||||
model_path = args.model_path
|
||||
path = maybe_download_model(model_path)
|
||||
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
ENCODER_PATH = os.path.join(path, "text_encoder")
|
||||
TOKENIZER_PATH = os.path.join(path, "tokenizer")
|
||||
print(ENCODER_PATH)
|
||||
text_encoder = text_encoder_loader.load(ENCODER_PATH, "text_encoder", fastvideo_args)
|
||||
tokenizer = tokenizer_loader.load(TOKENIZER_PATH, "tokenizer", fastvideo_args)
|
||||
|
||||
# text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
|
||||
# autocast_type = torch.float16 if args.model_type == "hunyuan" else torch.bfloat16
|
||||
# output_dir/validation/prompt_attention_mask
|
||||
# output_dir/validation/prompt_embed
|
||||
os.makedirs(os.path.join(args.output_dir, "validation"), exist_ok=True)
|
||||
os.makedirs(
|
||||
os.path.join(args.output_dir, "validation", "prompt_attention_mask"),
|
||||
exist_ok=True,
|
||||
)
|
||||
os.makedirs(os.path.join(args.output_dir, "validation", "prompt_embed"), exist_ok=True)
|
||||
|
||||
with open(args.validation_prompt_txt, "r", encoding="utf-8") as file:
|
||||
lines = file.readlines()
|
||||
prompts = [line.strip() for line in lines]
|
||||
for prompt in prompts:
|
||||
with torch.inference_mode():
|
||||
# with torch.autocast("cuda", dtype=autocast_type):
|
||||
text_inputs = tokenizer(prompt, **fastvideo_args.text_encoder_configs[0].tokenizer_kwargs).to(
|
||||
fastvideo_args.device)
|
||||
input_ids = text_inputs["input_ids"]
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
outputs = text_encoder(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
from fastvideo.v1.configs.pipelines.wan import t5_postprocess_text
|
||||
post_process_func = t5_postprocess_text
|
||||
prompt_embeds = post_process_func(outputs)
|
||||
prompt_attention_mask = attention_mask
|
||||
|
||||
file_name = prompt.split(".")[0]
|
||||
prompt_embed_path = os.path.join(args.output_dir, "validation", "prompt_embed", f"{file_name}.pt")
|
||||
prompt_attention_mask_path = os.path.join(
|
||||
args.output_dir,
|
||||
"validation",
|
||||
"prompt_attention_mask",
|
||||
f"{file_name}.pt",
|
||||
)
|
||||
torch.save(prompt_embeds[0], prompt_embed_path)
|
||||
torch.save(prompt_attention_mask[0], prompt_attention_mask_path)
|
||||
print(f"sample {file_name} saved")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--validation_prompt_txt", type=str)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -1,97 +0,0 @@
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.dataset.t2v_datasets import T2V_dataset
|
||||
from fastvideo.dataset.transform import CenterCropResizeVideo, Normalize255, TemporalRandomCrop
|
||||
|
||||
|
||||
def getdataset(args):
|
||||
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
|
||||
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
|
||||
resize_topcrop = [
|
||||
CenterCropResizeVideo((args.max_height, args.max_width), top_crop=True),
|
||||
]
|
||||
resize = [
|
||||
CenterCropResizeVideo((args.max_height, args.max_width)),
|
||||
]
|
||||
transform = transforms.Compose([
|
||||
# Normalize255(),
|
||||
*resize,
|
||||
])
|
||||
transform_topcrop = transforms.Compose([
|
||||
Normalize255(),
|
||||
*resize_topcrop,
|
||||
norm_fun,
|
||||
])
|
||||
# tokenizer = AutoTokenizer.from_pretrained("/storage/ongoing/new/Open-Sora-Plan/cache_dir/mt5-xxl", cache_dir=args.cache_dir)
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name, cache_dir=args.cache_dir)
|
||||
if args.dataset == "t2v":
|
||||
return T2V_dataset(
|
||||
args,
|
||||
transform=transform,
|
||||
temporal_sample=temporal_sample,
|
||||
tokenizer=tokenizer,
|
||||
transform_topcrop=transform_topcrop,
|
||||
)
|
||||
|
||||
raise NotImplementedError(args.dataset)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import random
|
||||
|
||||
from accelerate import Accelerator
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset.t2v_datasets import dataset_prog
|
||||
|
||||
args = type(
|
||||
"args",
|
||||
(),
|
||||
{
|
||||
"ae": "CausalVAEModel_4x8x8",
|
||||
"dataset": "t2v",
|
||||
"attention_mode": "xformers",
|
||||
"use_rope": True,
|
||||
"text_max_length": 300,
|
||||
"max_height": 320,
|
||||
"max_width": 240,
|
||||
"num_frames": 1,
|
||||
"use_image_num": 0,
|
||||
"interpolation_scale_t": 1,
|
||||
"interpolation_scale_h": 1,
|
||||
"interpolation_scale_w": 1,
|
||||
"cache_dir": "../cache_dir",
|
||||
"image_data": "/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
|
||||
"video_data": "1",
|
||||
"train_fps": 24,
|
||||
"drop_short_ratio": 1.0,
|
||||
"use_img_from_vid": False,
|
||||
"speed_factor": 1.0,
|
||||
"cfg": 0.1,
|
||||
"text_encoder_name": "google/mt5-xxl",
|
||||
"dataloader_num_workers": 10,
|
||||
},
|
||||
)
|
||||
accelerator = Accelerator()
|
||||
dataset = getdataset(args)
|
||||
num = len(dataset_prog.img_cap_list)
|
||||
zero = 0
|
||||
for idx in tqdm(range(num)):
|
||||
image_data = dataset_prog.img_cap_list[idx]
|
||||
caps = [i["cap"] if isinstance(i["cap"], list) else [i["cap"]] for i in image_data]
|
||||
try:
|
||||
caps = [[random.choice(i)] for i in caps]
|
||||
except Exception as e:
|
||||
print(e)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
print(image_data)
|
||||
zero += 1
|
||||
continue
|
||||
assert caps[0] is not None and len(caps[0]) > 0
|
||||
print(num, zero)
|
||||
import ipdb
|
||||
|
||||
ipdb.set_trace()
|
||||
print("end")
|
||||
@@ -1,118 +0,0 @@
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
|
||||
class LatentDataset(Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
json_path,
|
||||
num_latent_t,
|
||||
cfg_rate,
|
||||
):
|
||||
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
|
||||
self.json_path = json_path
|
||||
self.cfg_rate = cfg_rate
|
||||
self.datase_dir_path = os.path.dirname(json_path)
|
||||
self.video_dir = os.path.join(self.datase_dir_path, "video")
|
||||
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
|
||||
self.prompt_embed_dir = os.path.join(self.datase_dir_path, "prompt_embed")
|
||||
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path, "prompt_attention_mask")
|
||||
with open(self.json_path, "r") as f:
|
||||
self.data_anno = json.load(f)
|
||||
# json.load(f) already keeps the order
|
||||
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
|
||||
self.num_latent_t = num_latent_t
|
||||
# just zero embeddings [256, 4096]
|
||||
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
|
||||
# 256 zeros
|
||||
self.uncond_prompt_mask = torch.zeros(256).bool()
|
||||
self.lengths = [data_item["length"] if "length" in data_item else 1 for data_item in self.data_anno]
|
||||
|
||||
def __getitem__(self, idx):
|
||||
latent_file = self.data_anno[idx]["latent_path"]
|
||||
prompt_embed_file = self.data_anno[idx]["prompt_embed_path"]
|
||||
prompt_attention_mask_file = self.data_anno[idx]["prompt_attention_mask"]
|
||||
# load
|
||||
latent = torch.load(
|
||||
os.path.join(self.latent_dir, latent_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
latent = latent.squeeze(0)[:, -self.num_latent_t:]
|
||||
if random.random() < self.cfg_rate:
|
||||
prompt_embed = self.uncond_prompt_embed
|
||||
prompt_attention_mask = self.uncond_prompt_mask
|
||||
else:
|
||||
prompt_embed = torch.load(
|
||||
os.path.join(self.prompt_embed_dir, prompt_embed_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
prompt_attention_mask = torch.load(
|
||||
os.path.join(self.prompt_attention_mask_dir, prompt_attention_mask_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
return latent, prompt_embed, prompt_attention_mask
|
||||
|
||||
def __len__(self):
|
||||
return len(self.data_anno)
|
||||
|
||||
|
||||
def latent_collate_function(batch):
|
||||
# return latent, prompt, latent_attn_mask, text_attn_mask
|
||||
# latent_attn_mask: # b t h w
|
||||
# text_attn_mask: b 1 l
|
||||
# needs to check if the latent/prompt' size and apply padding & attn mask
|
||||
latents, prompt_embeds, prompt_attention_masks = zip(*batch)
|
||||
# calculate max shape
|
||||
max_t = max([latent.shape[1] for latent in latents])
|
||||
max_h = max([latent.shape[2] for latent in latents])
|
||||
max_w = max([latent.shape[3] for latent in latents])
|
||||
|
||||
# padding
|
||||
latents = [
|
||||
torch.nn.functional.pad(
|
||||
latent,
|
||||
(
|
||||
0,
|
||||
max_t - latent.shape[1],
|
||||
0,
|
||||
max_h - latent.shape[2],
|
||||
0,
|
||||
max_w - latent.shape[3],
|
||||
),
|
||||
) for latent in latents
|
||||
]
|
||||
# attn mask
|
||||
latent_attn_mask = torch.ones(len(latents), max_t, max_h, max_w)
|
||||
# set to 0 if padding
|
||||
for i, latent in enumerate(latents):
|
||||
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
|
||||
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
|
||||
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
|
||||
|
||||
prompt_embeds = torch.stack(prompt_embeds, dim=0)
|
||||
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
|
||||
latents = torch.stack(latents, dim=0)
|
||||
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt", num_latent_t=28)
|
||||
dataloader = torch.utils.data.DataLoader(dataset, batch_size=2, shuffle=False, collate_fn=latent_collate_function)
|
||||
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
|
||||
print(
|
||||
latent.shape,
|
||||
prompt_embed.shape,
|
||||
latent_attn_mask.shape,
|
||||
prompt_attention_mask.shape,
|
||||
)
|
||||
import pdb
|
||||
|
||||
pdb.set_trace()
|
||||
@@ -1,324 +0,0 @@
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import Counter
|
||||
from os.path import join as opj
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from fastvideo.utils.dataset_utils import DecordInit
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
|
||||
class SingletonMeta(type):
|
||||
_instances = {}
|
||||
|
||||
def __call__(cls, *args, **kwargs):
|
||||
if cls not in cls._instances:
|
||||
instance = super().__call__(*args, **kwargs)
|
||||
cls._instances[cls] = instance
|
||||
return cls._instances[cls]
|
||||
|
||||
|
||||
class DataSetProg(metaclass=SingletonMeta):
|
||||
|
||||
def __init__(self):
|
||||
self.cap_list = []
|
||||
self.elements = []
|
||||
self.num_workers = 1
|
||||
self.n_elements = 0
|
||||
self.worker_elements = dict()
|
||||
self.n_used_elements = dict()
|
||||
|
||||
def set_cap_list(self, num_workers, cap_list, n_elements):
|
||||
self.num_workers = num_workers
|
||||
self.cap_list = cap_list
|
||||
self.n_elements = n_elements
|
||||
self.elements = list(range(n_elements))
|
||||
random.shuffle(self.elements)
|
||||
print(f"n_elements: {len(self.elements)}", flush=True)
|
||||
|
||||
for i in range(self.num_workers):
|
||||
self.n_used_elements[i] = 0
|
||||
per_worker = int(math.ceil(len(self.elements) / float(self.num_workers)))
|
||||
start = i * per_worker
|
||||
end = min(start + per_worker, len(self.elements))
|
||||
self.worker_elements[i] = self.elements[start:end]
|
||||
|
||||
def get_item(self, work_info):
|
||||
if work_info is None:
|
||||
worker_id = 0
|
||||
else:
|
||||
worker_id = work_info.id
|
||||
|
||||
idx = self.worker_elements[worker_id][self.n_used_elements[worker_id] % len(self.worker_elements[worker_id])]
|
||||
self.n_used_elements[worker_id] += 1
|
||||
return idx
|
||||
|
||||
|
||||
dataset_prog = DataSetProg()
|
||||
|
||||
|
||||
def filter_resolution(h, w, max_h_div_w_ratio=17 / 16, min_h_div_w_ratio=8 / 16):
|
||||
if h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
class T2V_dataset(Dataset):
|
||||
|
||||
def __init__(self, args, transform, temporal_sample, tokenizer, transform_topcrop):
|
||||
self.data = args.data_merge_path
|
||||
self.num_frames = args.num_frames
|
||||
self.train_fps = args.train_fps
|
||||
self.use_image_num = args.use_image_num
|
||||
self.transform = transform
|
||||
self.transform_topcrop = transform_topcrop
|
||||
self.temporal_sample = temporal_sample
|
||||
self.tokenizer = tokenizer
|
||||
self.text_max_length = args.text_max_length
|
||||
self.cfg = args.cfg
|
||||
self.speed_factor = args.speed_factor
|
||||
self.max_height = args.max_height
|
||||
self.max_width = args.max_width
|
||||
self.drop_short_ratio = args.drop_short_ratio
|
||||
assert self.speed_factor >= 1
|
||||
self.v_decoder = DecordInit()
|
||||
self.video_length_tolerance_range = args.video_length_tolerance_range
|
||||
self.support_Chinese = True
|
||||
if "mt5" not in args.text_encoder_name:
|
||||
self.support_Chinese = False
|
||||
|
||||
cap_list = self.get_cap_list()
|
||||
|
||||
assert len(cap_list) > 0
|
||||
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
|
||||
self.lengths = self.sample_num_frames
|
||||
|
||||
n_elements = len(cap_list)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list, n_elements)
|
||||
|
||||
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
|
||||
|
||||
def set_checkpoint(self, n_used_elements):
|
||||
for i in range(len(dataset_prog.n_used_elements)):
|
||||
dataset_prog.n_used_elements[i] = n_used_elements
|
||||
|
||||
def __len__(self):
|
||||
return dataset_prog.n_elements
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
data = self.get_data(idx)
|
||||
return data
|
||||
|
||||
def get_data(self, idx):
|
||||
path = dataset_prog.cap_list[idx]["path"]
|
||||
if path.endswith(".mp4"):
|
||||
return self.get_video(idx)
|
||||
else:
|
||||
return self.get_image(idx)
|
||||
|
||||
def get_video(self, idx):
|
||||
video_path = dataset_prog.cap_list[idx]["path"]
|
||||
assert os.path.exists(video_path), f"file {video_path} do not exist!"
|
||||
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(video_path, output_format="TCHW")
|
||||
video = torchvision_video[frame_indices]
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
assert video.dtype == torch.uint8
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
|
||||
text = dataset_prog.cap_list[idx]["cap"]
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
|
||||
text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"]
|
||||
cond_mask = text_tokens_and_mask["attention_mask"]
|
||||
return dict(
|
||||
pixel_values=video,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=video_path,
|
||||
)
|
||||
|
||||
def get_image(self, idx):
|
||||
image_data = dataset_prog.cap_list[idx] # [{'path': path, 'cap': cap}, ...]
|
||||
|
||||
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
|
||||
image = torch.from_numpy(np.array(image)) # [h, w, c]
|
||||
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
|
||||
# for i in image:
|
||||
# h, w = i.shape[-2:]
|
||||
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
|
||||
|
||||
image = (self.transform_topcrop(image) if "human_images" in image_data["path"] else self.transform(image)
|
||||
) # [1 C H W] -> num_img [1 C H W]
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
|
||||
image = image.float() / 127.5 - 1.0
|
||||
|
||||
caps = (image_data["cap"] if isinstance(image_data["cap"], list) else [image_data["cap"]])
|
||||
caps = [random.choice(caps)]
|
||||
text = caps
|
||||
input_ids, cond_mask = [], []
|
||||
text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"] # 1, l
|
||||
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
|
||||
return dict(
|
||||
pixel_values=image,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=image_data["path"],
|
||||
)
|
||||
|
||||
def define_frame_index(self, cap_list):
|
||||
new_cap_list = []
|
||||
sample_num_frames = []
|
||||
cnt_too_long = 0
|
||||
cnt_too_short = 0
|
||||
cnt_no_cap = 0
|
||||
cnt_no_resolution = 0
|
||||
cnt_resolution_mismatch = 0
|
||||
cnt_movie = 0
|
||||
cnt_img = 0
|
||||
for i in cap_list:
|
||||
path = i["path"]
|
||||
cap = i.get("cap", None)
|
||||
# ======no caption=====
|
||||
if cap is None:
|
||||
cnt_no_cap += 1
|
||||
continue
|
||||
if path.endswith(".mp4"):
|
||||
# ======no fps and duration=====
|
||||
duration = i.get("duration", None)
|
||||
fps = i.get("fps", None)
|
||||
if fps is None or duration is None:
|
||||
continue
|
||||
|
||||
# ======resolution mismatch=====
|
||||
resolution = i.get("resolution", None)
|
||||
if resolution is None:
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
else:
|
||||
if (resolution.get("height", None) is None or resolution.get("width", None) is None):
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
height, width = i["resolution"]["height"], i["resolution"]["width"]
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
is_pick = filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
if not is_pick:
|
||||
print("resolution mismatch")
|
||||
cnt_resolution_mismatch += 1
|
||||
continue
|
||||
|
||||
# import ipdb;ipdb.set_trace()
|
||||
i["num_frames"] = math.ceil(fps * duration)
|
||||
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
|
||||
if i["num_frames"] / fps > self.video_length_tolerance_range * (
|
||||
self.num_frames / self.train_fps *
|
||||
self.speed_factor): # too long video is not suitable for this training stage (self.num_frames)
|
||||
cnt_too_long += 1
|
||||
continue
|
||||
|
||||
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
|
||||
frame_interval = fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, i["num_frames"], frame_interval).astype(int)
|
||||
|
||||
# comment out it to enable dynamic frames training
|
||||
if (len(frame_indices) < self.num_frames and random.random() < self.drop_short_ratio):
|
||||
cnt_too_short += 1
|
||||
continue
|
||||
|
||||
# too long video will be temporal-crop randomly
|
||||
if len(frame_indices) > self.num_frames:
|
||||
begin_index, end_index = self.temporal_sample(len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index:end_index]
|
||||
# frame_indices = frame_indices[:self.num_frames] # head crop
|
||||
i["sample_frame_index"] = frame_indices.tolist()
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = len(i["sample_frame_index"]) # will use in dataloader(group sampler)
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
elif path.endswith(".jpg"): # image
|
||||
cnt_img += 1
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = 1
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
else:
|
||||
raise NameError(
|
||||
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image")
|
||||
# import ipdb;ipdb.set_trace()
|
||||
main_print(
|
||||
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
|
||||
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
|
||||
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
|
||||
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}")
|
||||
return new_cap_list, sample_num_frames
|
||||
|
||||
def decord_read(self, path, frame_indices):
|
||||
decord_vr = self.v_decoder(path)
|
||||
video_data = decord_vr.get_batch(frame_indices).asnumpy()
|
||||
video_data = torch.from_numpy(video_data)
|
||||
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
|
||||
return video_data
|
||||
|
||||
def read_jsons(self, data):
|
||||
cap_lists = []
|
||||
with open(data, "r") as f:
|
||||
folder_anno = [i.strip().split(",") for i in f.readlines() if len(i.strip()) > 0]
|
||||
print(folder_anno)
|
||||
for folder, anno in folder_anno:
|
||||
with open(anno, "r") as f:
|
||||
sub_list = json.load(f)
|
||||
for i in range(len(sub_list)):
|
||||
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
|
||||
cap_lists += sub_list
|
||||
return cap_lists
|
||||
|
||||
def get_cap_list(self):
|
||||
cap_lists = self.read_jsons(self.data)
|
||||
return cap_lists
|
||||
@@ -1,608 +0,0 @@
|
||||
import numbers
|
||||
import random
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def _is_tensor_video_clip(clip):
|
||||
if not torch.is_tensor(clip):
|
||||
raise TypeError("clip should be Tensor. Got %s" % type(clip))
|
||||
|
||||
if not clip.ndimension() == 4:
|
||||
raise ValueError("clip should be 4D. Got %dD" % clip.dim())
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def center_crop_arr(pil_image, image_size):
|
||||
"""
|
||||
Center cropping implementation from ADM.
|
||||
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
|
||||
"""
|
||||
while min(*pil_image.size) >= 2 * image_size:
|
||||
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size), resample=Image.BOX)
|
||||
|
||||
scale = image_size / min(*pil_image.size)
|
||||
pil_image = pil_image.resize(tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC)
|
||||
|
||||
arr = np.array(pil_image)
|
||||
crop_y = (arr.shape[0] - image_size) // 2
|
||||
crop_x = (arr.shape[1] - image_size) // 2
|
||||
return Image.fromarray(arr[crop_y:crop_y + image_size, crop_x:crop_x + image_size])
|
||||
|
||||
|
||||
def crop(clip, i, j, h, w):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
"""
|
||||
if len(clip.size()) != 4:
|
||||
raise ValueError("clip should be a 4D tensor")
|
||||
return clip[..., i:i + h, j:j + w]
|
||||
|
||||
|
||||
def resize(clip, target_size, interpolation_mode):
|
||||
if len(target_size) != 2:
|
||||
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
|
||||
return torch.nn.functional.interpolate(
|
||||
clip,
|
||||
size=target_size,
|
||||
mode=interpolation_mode,
|
||||
align_corners=True,
|
||||
antialias=True,
|
||||
)
|
||||
|
||||
|
||||
def resize_scale(clip, target_size, interpolation_mode):
|
||||
if len(target_size) != 2:
|
||||
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
|
||||
H, W = clip.size(-2), clip.size(-1)
|
||||
scale_ = target_size[0] / min(H, W)
|
||||
return torch.nn.functional.interpolate(
|
||||
clip,
|
||||
scale_factor=scale_,
|
||||
mode=interpolation_mode,
|
||||
align_corners=True,
|
||||
antialias=True,
|
||||
)
|
||||
|
||||
|
||||
def resized_crop(clip, i, j, h, w, size, interpolation_mode="bilinear"):
|
||||
"""
|
||||
Do spatial cropping and resizing to the video clip
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
i (int): i in (i,j) i.e coordinates of the upper left corner.
|
||||
j (int): j in (i,j) i.e coordinates of the upper left corner.
|
||||
h (int): Height of the cropped region.
|
||||
w (int): Width of the cropped region.
|
||||
size (tuple(int, int)): height and width of resized clip
|
||||
Returns:
|
||||
clip (torch.tensor): Resized and cropped clip. Size is (T, C, H, W)
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
clip = crop(clip, i, j, h, w)
|
||||
clip = resize(clip, size, interpolation_mode)
|
||||
return clip
|
||||
|
||||
|
||||
def center_crop(clip, crop_size):
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
th, tw = crop_size
|
||||
if h < th or w < tw:
|
||||
raise ValueError("height and width must be no smaller than crop_size")
|
||||
|
||||
i = int(round((h - th) / 2.0))
|
||||
j = int(round((w - tw) / 2.0))
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
def center_crop_using_short_edge(clip):
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
if h < w:
|
||||
th, tw = h, h
|
||||
i = 0
|
||||
j = int(round((w - tw) / 2.0))
|
||||
else:
|
||||
th, tw = w, w
|
||||
i = int(round((h - th) / 2.0))
|
||||
j = 0
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
def center_crop_th_tw(clip, th, tw, top_crop):
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
|
||||
# import ipdb;ipdb.set_trace()
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
tr = th / tw
|
||||
if h / w > tr:
|
||||
new_h = int(w * tr)
|
||||
new_w = w
|
||||
else:
|
||||
new_h = h
|
||||
new_w = int(h / tr)
|
||||
|
||||
i = 0 if top_crop else int(round((h - new_h) / 2.0))
|
||||
j = int(round((w - new_w) / 2.0))
|
||||
return crop(clip, i, j, new_h, new_w)
|
||||
|
||||
|
||||
def random_shift_crop(clip):
|
||||
"""
|
||||
Slide along the long edge, with the short edge as crop size
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
|
||||
if h <= w:
|
||||
short_edge = h
|
||||
else:
|
||||
short_edge = w
|
||||
|
||||
th, tw = short_edge, short_edge
|
||||
|
||||
i = torch.randint(0, h - th + 1, size=(1, )).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1, )).item()
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
def normalize_video(clip):
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
permute the dimensions of clip tensor
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
_is_tensor_video_clip(clip)
|
||||
if not clip.dtype == torch.uint8:
|
||||
raise TypeError("clip tensor should have data type uint8. Got %s" % str(clip.dtype))
|
||||
# return clip.float().permute(3, 0, 1, 2) / 255.0
|
||||
return clip.float() / 255.0
|
||||
|
||||
|
||||
def normalize(clip, mean, std, inplace=False):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
|
||||
mean (tuple): pixel RGB mean. Size is (3)
|
||||
std (tuple): pixel standard deviation. Size is (3)
|
||||
Returns:
|
||||
normalized clip (torch.tensor): Size is (T, C, H, W)
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
if not inplace:
|
||||
clip = clip.clone()
|
||||
mean = torch.as_tensor(mean, dtype=clip.dtype, device=clip.device)
|
||||
# print(mean)
|
||||
std = torch.as_tensor(std, dtype=clip.dtype, device=clip.device)
|
||||
clip.sub_(mean[:, None, None, None]).div_(std[:, None, None, None])
|
||||
return clip
|
||||
|
||||
|
||||
def hflip(clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
|
||||
Returns:
|
||||
flipped clip (torch.tensor): Size is (T, C, H, W)
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
return clip.flip(-1)
|
||||
|
||||
|
||||
class RandomCropVideo:
|
||||
|
||||
def __init__(self, size):
|
||||
if isinstance(size, numbers.Number):
|
||||
self.size = (int(size), int(size))
|
||||
else:
|
||||
self.size = size
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: randomly cropped video clip.
|
||||
size is (T, C, OH, OW)
|
||||
"""
|
||||
i, j, h, w = self.get_params(clip)
|
||||
return crop(clip, i, j, h, w)
|
||||
|
||||
def get_params(self, clip):
|
||||
h, w = clip.shape[-2:]
|
||||
th, tw = self.size
|
||||
|
||||
if h < th or w < tw:
|
||||
raise ValueError(f"Required crop size {(th, tw)} is larger than input image size {(h, w)}")
|
||||
|
||||
if w == tw and h == th:
|
||||
return 0, 0, h, w
|
||||
|
||||
i = torch.randint(0, h - th + 1, size=(1, )).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1, )).item()
|
||||
|
||||
return i, j, th, tw
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size})"
|
||||
|
||||
|
||||
class SpatialStrideCropVideo:
|
||||
|
||||
def __init__(self, stride):
|
||||
self.stride = stride
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: cropped video clip by stride.
|
||||
size is (T, C, OH, OW)
|
||||
"""
|
||||
i, j, h, w = self.get_params(clip)
|
||||
return crop(clip, i, j, h, w)
|
||||
|
||||
def get_params(self, clip):
|
||||
h, w = clip.shape[-2:]
|
||||
|
||||
th, tw = h // self.stride * self.stride, w // self.stride * self.stride
|
||||
|
||||
return 0, 0, th, tw # from top-left
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size})"
|
||||
|
||||
|
||||
class LongSideResizeVideo:
|
||||
"""
|
||||
First use the long side,
|
||||
then resize to the specified size
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
skip_low_resolution=False,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
self.size = size
|
||||
self.skip_low_resolution = skip_low_resolution
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: scale resized video clip.
|
||||
size is (T, C, 512, *) or (T, C, *, 512)
|
||||
"""
|
||||
_, _, h, w = clip.shape
|
||||
if self.skip_low_resolution and max(h, w) <= self.size:
|
||||
return clip
|
||||
if h > w:
|
||||
w = int(w * self.size / h)
|
||||
h = self.size
|
||||
else:
|
||||
h = int(h * self.size / w)
|
||||
w = self.size
|
||||
resize_clip = resize(clip, target_size=(h, w), interpolation_mode=self.interpolation_mode)
|
||||
return resize_clip
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class CenterCropResizeVideo:
|
||||
"""
|
||||
First use the short side for cropping length,
|
||||
center crop video, then resize to the specified size
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
top_crop=False,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
self.top_crop = top_crop
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: scale resized / center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
# clip_center_crop = center_crop_using_short_edge(clip)
|
||||
clip_center_crop = center_crop_th_tw(clip, self.size[0], self.size[1], top_crop=self.top_crop)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
clip_center_crop_resize = resize(
|
||||
clip_center_crop,
|
||||
target_size=self.size,
|
||||
interpolation_mode=self.interpolation_mode,
|
||||
)
|
||||
return clip_center_crop_resize
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class UCFCenterCropVideo:
|
||||
"""
|
||||
First scale to the specified size in equal proportion to the short edge,
|
||||
then center cropping
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: scale resized / center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
clip_resize = resize_scale(clip=clip, target_size=self.size, interpolation_mode=self.interpolation_mode)
|
||||
clip_center_crop = center_crop(clip_resize, self.size)
|
||||
return clip_center_crop
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class KineticsRandomCropResizeVideo:
|
||||
"""
|
||||
Slide along the long edge, with the short edge as crop size. And resie to the desired size.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
clip_random_crop = random_shift_crop(clip)
|
||||
clip_resize = resize(clip_random_crop, self.size, self.interpolation_mode)
|
||||
return clip_resize
|
||||
|
||||
|
||||
class CenterCropVideo:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
clip_center_crop = center_crop(clip, self.size)
|
||||
return clip_center_crop
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class Normalize:
|
||||
"""
|
||||
Normalize the video clip by mean subtraction and division by standard deviation
|
||||
Args:
|
||||
mean (3-tuple): pixel RGB mean
|
||||
std (3-tuple): pixel RGB standard deviation
|
||||
inplace (boolean): whether do in-place normalization
|
||||
"""
|
||||
|
||||
def __init__(self, mean, std, inplace=False):
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
self.inplace = inplace
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): video clip must be normalized. Size is (C, T, H, W)
|
||||
"""
|
||||
return normalize(clip, self.mean, self.std, self.inplace)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(mean={self.mean}, std={self.std}, inplace={self.inplace})"
|
||||
|
||||
|
||||
class Normalize255:
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
return normalize_video(clip)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return self.__class__.__name__
|
||||
|
||||
|
||||
class RandomHorizontalFlipVideo:
|
||||
"""
|
||||
Flip the video clip along the horizontal direction with a given probability
|
||||
Args:
|
||||
p (float): probability of the clip being flipped. Default value is 0.5
|
||||
"""
|
||||
|
||||
def __init__(self, p=0.5):
|
||||
self.p = p
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor): Size is (T, C, H, W)
|
||||
"""
|
||||
if random.random() < self.p:
|
||||
clip = hflip(clip)
|
||||
return clip
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(p={self.p})"
|
||||
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# --------------------- Sampling ---------------------------
|
||||
# ------------------------------------------------------------
|
||||
class TemporalRandomCrop(object):
|
||||
"""Temporally crop the given frame indices at a random location.
|
||||
|
||||
Args:
|
||||
size (int): Desired length of frames will be seen in the model.
|
||||
"""
|
||||
|
||||
def __init__(self, size):
|
||||
self.size = size
|
||||
|
||||
def __call__(self, total_frames):
|
||||
rand_end = max(0, total_frames - self.size - 1)
|
||||
begin_index = random.randint(0, rand_end)
|
||||
end_index = min(begin_index + self.size, total_frames)
|
||||
return begin_index, end_index
|
||||
|
||||
|
||||
class DynamicSampleDuration(object):
|
||||
"""Temporally crop the given frame indices at a random location.
|
||||
|
||||
Args:
|
||||
size (int): Desired length of frames will be seen in the model.
|
||||
"""
|
||||
|
||||
def __init__(self, t_stride, extra_1):
|
||||
self.t_stride = t_stride
|
||||
self.extra_1 = extra_1
|
||||
|
||||
def __call__(self, t, h, w):
|
||||
if self.extra_1:
|
||||
t = t - 1
|
||||
truncate_t_list = list(range(t + 1))[t // 2:][::self.t_stride] # need half at least
|
||||
truncate_t = random.choice(truncate_t_list)
|
||||
if self.extra_1:
|
||||
truncate_t = truncate_t + 1
|
||||
return 0, truncate_t
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torchvision.io as io
|
||||
from torchvision import transforms
|
||||
from torchvision.utils import save_image
|
||||
|
||||
vframes, aframes, info = io.read_video(filename="./v_Archery_g01_c03.avi", pts_unit="sec", output_format="TCHW")
|
||||
|
||||
trans = transforms.Compose([
|
||||
Normalize255(),
|
||||
RandomHorizontalFlipVideo(),
|
||||
UCFCenterCropVideo(512),
|
||||
# NormalizeVideo(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
])
|
||||
|
||||
target_video_len = 32
|
||||
frame_interval = 1
|
||||
total_frames = len(vframes)
|
||||
print(total_frames)
|
||||
|
||||
temporal_sample = TemporalRandomCrop(target_video_len * frame_interval)
|
||||
|
||||
# Sampling video frames
|
||||
start_frame_ind, end_frame_ind = temporal_sample(total_frames)
|
||||
# print(start_frame_ind)
|
||||
# print(end_frame_ind)
|
||||
assert end_frame_ind - start_frame_ind >= target_video_len
|
||||
frame_indice = np.linspace(start_frame_ind, end_frame_ind - 1, target_video_len, dtype=int)
|
||||
print(frame_indice)
|
||||
|
||||
select_vframes = vframes[frame_indice]
|
||||
print(select_vframes.shape)
|
||||
print(select_vframes.dtype)
|
||||
|
||||
select_vframes_trans = trans(select_vframes)
|
||||
print(select_vframes_trans.shape)
|
||||
print(select_vframes_trans.dtype)
|
||||
|
||||
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) * 255).to(dtype=torch.uint8)
|
||||
print(select_vframes_trans_int.dtype)
|
||||
print(select_vframes_trans_int.permute(0, 2, 3, 1).shape)
|
||||
|
||||
io.write_video("./test.avi", select_vframes_trans_int.permute(0, 2, 3, 1), fps=8)
|
||||
|
||||
for i in range(target_video_len):
|
||||
save_image(
|
||||
select_vframes_trans[i],
|
||||
os.path.join("./test000", "%04d.png" % i),
|
||||
normalize=True,
|
||||
value_range=(-1, 1),
|
||||
)
|
||||
+47
-74
@@ -12,7 +12,6 @@ import torch.distributed as dist
|
||||
import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
from diffusers.optimization import get_scheduler
|
||||
from diffusers.utils import check_min_version
|
||||
from peft import LoraConfig
|
||||
@@ -24,7 +23,7 @@ from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from fastvideo.utils.latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint, save_lora_checkpoint)
|
||||
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
|
||||
@@ -124,21 +123,13 @@ def distill_one_step(
|
||||
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
|
||||
# Predict the noise residual
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
if args.model_type == "wan":
|
||||
teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"return_dict": True,
|
||||
}
|
||||
else:
|
||||
teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
if hunyuan_teacher_disable_cfg:
|
||||
teacher_kwargs["guidance"] = torch.tensor([1000.0],
|
||||
device=noisy_model_input.device,
|
||||
@@ -150,70 +141,47 @@ def distill_one_step(
|
||||
with torch.no_grad():
|
||||
w = distill_cfg
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
if args.model_type == "wan":
|
||||
cond_teacher_kwargs ={
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"return_dict": True,
|
||||
}
|
||||
else:
|
||||
cond_teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
cond_teacher_output = teacher_transformer(**cond_teacher_kwargs)[0].float()
|
||||
cond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
if not_apply_cfg_solver:
|
||||
uncond_teacher_output = cond_teacher_output
|
||||
else:
|
||||
# Get teacher model prediction on noisy_latents and unconditional embedding
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
if args.model_type == "wan":
|
||||
uncond_teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states":uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
"timestep": timesteps,
|
||||
"return_dict": True,
|
||||
}
|
||||
else:
|
||||
uncond_teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states":uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
"return_dict": False,
|
||||
}
|
||||
|
||||
uncond_teacher_output = teacher_transformer(**uncond_teacher_kwargs)[0].float()
|
||||
|
||||
uncond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
timesteps,
|
||||
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
teacher_output = uncond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
|
||||
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
|
||||
|
||||
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
|
||||
with torch.no_grad():
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
if args.model_type == "wan":
|
||||
target_pred_kwargs = {
|
||||
"hidden_states": x_prev.float(),
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep":timesteps_prev,
|
||||
"return_dict":True,
|
||||
}
|
||||
else:
|
||||
target_pred_kwargs = {
|
||||
"hidden_states": x_prev.float(),
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep":timesteps_prev,
|
||||
"encoder_attention_mask":encoder_attention_mask,
|
||||
"return_dict":False,
|
||||
}
|
||||
if ema_transformer is not None:
|
||||
target_pred = ema_transformer(**target_pred_kwargs)[0]
|
||||
target_pred = ema_transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0]
|
||||
else:
|
||||
target_pred = transformer(**target_pred_kwargs)[0]
|
||||
target_pred = transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
target, end_index = solver.euler_style_multiphase_pred(x_prev, target_pred, index, multiphase, True)
|
||||
|
||||
@@ -274,7 +242,7 @@ def main(args):
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank == 0 and args.output_dir is not None:
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
@@ -351,9 +319,7 @@ def main(args):
|
||||
teacher_transformer.requires_grad_(False)
|
||||
if args.use_ema:
|
||||
ema_transformer.requires_grad_(False)
|
||||
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
|
||||
if args.scheduler_type == "pcm_linear_quadratic":
|
||||
linear_steps = int(noise_scheduler.config.num_train_timesteps * args.linear_range)
|
||||
sigmas = linear_quadratic_schedule(
|
||||
@@ -425,7 +391,7 @@ def main(args):
|
||||
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
if rank == 0:
|
||||
if rank <= 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
@@ -527,7 +493,7 @@ def main(args):
|
||||
"phases": num_phases,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank == 0:
|
||||
if rank <= 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss": loss,
|
||||
@@ -671,6 +637,13 @@ if __name__ == "__main__":
|
||||
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
|
||||
@@ -7,7 +7,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
# from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
@@ -38,6 +38,7 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
linear_range=0.5,
|
||||
):
|
||||
if linear_quadratic:
|
||||
raise NotImplementedError("Linear quadratic schedule is not implemented")
|
||||
linear_steps = int(num_train_timesteps * linear_range)
|
||||
sigmas = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
|
||||
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
|
||||
|
||||
@@ -23,7 +23,7 @@ from tqdm.auto import tqdm
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.distill.discriminator import Discriminator
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from fastvideo.utils.latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
from fastvideo.utils.checkpoint import (resume_lora_optimizer, resume_training_generator_discriminator, save_checkpoint,
|
||||
save_lora_checkpoint)
|
||||
@@ -296,7 +296,7 @@ def main(args):
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank == 0 and args.output_dir is not None:
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
@@ -462,7 +462,7 @@ def main(args):
|
||||
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
if rank == 0:
|
||||
if rank <= 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
@@ -559,7 +559,7 @@ def main(args):
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank == 0:
|
||||
if rank <= 0:
|
||||
wandb.log(
|
||||
{
|
||||
"generator_loss": generator_loss,
|
||||
@@ -693,6 +693,13 @@ if __name__ == "__main__":
|
||||
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
|
||||
@@ -0,0 +1,870 @@
|
||||
# !/bin/python3
|
||||
# isort: skip_file
|
||||
import argparse
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from collections import deque
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from diffusers.optimization import get_scheduler
|
||||
from diffusers.utils import check_min_version
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.fsdp import ShardingStrategy
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset,
|
||||
latent_collate_function)
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from fastvideo.utils.checkpoint import (save_checkpoint, save_lora_checkpoint)
|
||||
from fastvideo.utils.communications import (broadcast,
|
||||
sp_parallel_dataloader_wrapper)
|
||||
from fastvideo.utils.dataset_utils import LengthGroupedSampler
|
||||
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group,
|
||||
get_sequence_parallel_state)
|
||||
from fastvideo.utils.validation import log_validation
|
||||
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.v1.models.loader.component_loader import TransformerLoader, SchedulerLoader
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.31.0")
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
BASE_MODEL_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
SCHEDULER_PATH = os.path.join(MODEL_PATH, "scheduler")
|
||||
|
||||
|
||||
def reshard_fsdp(model):
|
||||
for m in FSDP.fsdp_modules(model):
|
||||
if m._has_params and m.sharding_strategy is not ShardingStrategy.NO_SHARD:
|
||||
torch.distributed.fsdp._runtime_utils._reshard(m, m._handle, True)
|
||||
|
||||
|
||||
def get_norm(model_pred, norms, gradient_accumulation_steps):
|
||||
fro_norm = (
|
||||
torch.linalg.matrix_norm(model_pred, ord="fro") / # codespell:ignore
|
||||
gradient_accumulation_steps)
|
||||
largest_singular_value = (torch.linalg.matrix_norm(model_pred, ord=2) /
|
||||
gradient_accumulation_steps)
|
||||
absolute_mean = torch.mean(
|
||||
torch.abs(model_pred)) / gradient_accumulation_steps
|
||||
absolute_max = torch.max(
|
||||
torch.abs(model_pred)) / gradient_accumulation_steps
|
||||
dist.all_reduce(fro_norm, op=dist.ReduceOp.AVG)
|
||||
dist.all_reduce(largest_singular_value, op=dist.ReduceOp.AVG)
|
||||
dist.all_reduce(absolute_mean, op=dist.ReduceOp.AVG)
|
||||
norms["fro"] += torch.mean(fro_norm).item() # codespell:ignore
|
||||
norms["largest singular value"] += torch.mean(largest_singular_value).item()
|
||||
norms["absolute mean"] += absolute_mean.item()
|
||||
norms["absolute max"] += absolute_max.item()
|
||||
|
||||
|
||||
def distill_one_step(
|
||||
transformer,
|
||||
model_type,
|
||||
teacher_transformer,
|
||||
ema_transformer,
|
||||
optimizer,
|
||||
lr_scheduler,
|
||||
loader,
|
||||
noise_scheduler,
|
||||
solver,
|
||||
noise_random_generator,
|
||||
gradient_accumulation_steps,
|
||||
sp_size,
|
||||
max_grad_norm,
|
||||
uncond_prompt_embed,
|
||||
uncond_prompt_mask,
|
||||
num_euler_timesteps,
|
||||
multiphase,
|
||||
not_apply_cfg_solver,
|
||||
distill_cfg,
|
||||
ema_decay,
|
||||
pred_decay_weight,
|
||||
pred_decay_type,
|
||||
hunyuan_teacher_disable_cfg,
|
||||
):
|
||||
total_loss = 0.0
|
||||
optimizer.zero_grad()
|
||||
model_pred_norm = {
|
||||
"fro": 0.0, # codespell:ignore
|
||||
"largest singular value": 0.0,
|
||||
"absolute mean": 0.0,
|
||||
"absolute max": 0.0,
|
||||
}
|
||||
for _ in range(gradient_accumulation_steps):
|
||||
(
|
||||
latents,
|
||||
encoder_hidden_states,
|
||||
latents_attention_mask,
|
||||
encoder_attention_mask,
|
||||
) = next(loader)
|
||||
# model_input = normalize_dit_input(model_type, latents)
|
||||
model_input = latents
|
||||
noise = torch.randn_like(model_input)
|
||||
bsz = model_input.shape[0]
|
||||
index = torch.randint(0,
|
||||
num_euler_timesteps, (bsz, ),
|
||||
device=model_input.device).long()
|
||||
if sp_size > 1:
|
||||
broadcast(index)
|
||||
# Add noise according to flow matching.
|
||||
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
|
||||
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
|
||||
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index,
|
||||
model_input.shape)
|
||||
|
||||
timesteps = (sigmas *
|
||||
noise_scheduler.config.num_train_timesteps).view(-1)
|
||||
# if squeeze to [], unsqueeze to [1]
|
||||
|
||||
timesteps_prev = (sigmas_prev *
|
||||
noise_scheduler.config.num_train_timesteps).view(-1)
|
||||
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
|
||||
noisy_model_input = noisy_model_input.to(torch.bfloat16)
|
||||
|
||||
forward_batch = ForwardBatch(data_type="video", enable_teacache=False)
|
||||
# Predict the noise residual
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
if hunyuan_teacher_disable_cfg:
|
||||
teacher_kwargs["guidance"] = torch.tensor(
|
||||
[1000.0],
|
||||
device=noisy_model_input.device,
|
||||
dtype=torch.bfloat16)
|
||||
with set_forward_context(current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch):
|
||||
with torch.autograd.graph.save_on_cpu(pin_memory=True):
|
||||
model_pred = transformer(**teacher_kwargs)
|
||||
|
||||
# if accelerator.is_main_process:
|
||||
model_pred, end_index = solver.euler_style_multiphase_pred(
|
||||
noisy_model_input, model_pred, index, multiphase)
|
||||
with torch.no_grad():
|
||||
w = distill_cfg
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
with set_forward_context(current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch):
|
||||
cond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
).float()
|
||||
if not_apply_cfg_solver:
|
||||
uncond_teacher_output = cond_teacher_output
|
||||
else:
|
||||
# Get teacher model prediction on noisy_latents and unconditional embedding
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
with set_forward_context(current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch):
|
||||
uncond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
uncond_prompt_embed.unsqueeze(0).expand(
|
||||
bsz, -1, -1),
|
||||
timesteps,
|
||||
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
return_dict=False,
|
||||
).float()
|
||||
teacher_output = uncond_teacher_output + w * (cond_teacher_output -
|
||||
uncond_teacher_output)
|
||||
x_prev = solver.euler_step(noisy_model_input, teacher_output,
|
||||
index).to(torch.bfloat16)
|
||||
|
||||
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
|
||||
with torch.no_grad():
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
if ema_transformer is not None:
|
||||
target_pred = ema_transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0]
|
||||
else:
|
||||
with set_forward_context(current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch):
|
||||
with torch.autograd.graph.save_on_cpu(pin_memory=True):
|
||||
target_pred = transformer(
|
||||
x_prev,
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)
|
||||
|
||||
target, end_index = solver.euler_style_multiphase_pred(
|
||||
x_prev, target_pred, index, multiphase, True)
|
||||
|
||||
huber_c = 0.001
|
||||
# loss = loss.mean()
|
||||
loss = (torch.mean(
|
||||
torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) -
|
||||
huber_c) / gradient_accumulation_steps)
|
||||
if pred_decay_weight > 0:
|
||||
if pred_decay_type == "l1":
|
||||
pred_decay_loss = (
|
||||
torch.mean(torch.sqrt(model_pred.float()**2)) *
|
||||
pred_decay_weight / gradient_accumulation_steps)
|
||||
loss += pred_decay_loss
|
||||
elif pred_decay_type == "l2":
|
||||
# essnetially k2?
|
||||
pred_decay_loss = (torch.mean(model_pred.float()**2) *
|
||||
pred_decay_weight /
|
||||
gradient_accumulation_steps)
|
||||
loss += pred_decay_loss
|
||||
else:
|
||||
assert NotImplementedError("pred_decay_type is not implemented")
|
||||
|
||||
# calculate model_pred norm and mean
|
||||
get_norm(model_pred.detach().float(), model_pred_norm,
|
||||
gradient_accumulation_steps)
|
||||
loss.backward()
|
||||
|
||||
avg_loss = loss.detach().clone()
|
||||
dist.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
|
||||
total_loss += avg_loss.item()
|
||||
|
||||
# update ema
|
||||
if ema_transformer is not None:
|
||||
reshard_fsdp(ema_transformer)
|
||||
for p_averaged, p_model in zip(ema_transformer.parameters(),
|
||||
transformer.parameters()):
|
||||
with torch.no_grad():
|
||||
p_averaged.copy_(
|
||||
torch.lerp(p_averaged.detach(), p_model.detach(),
|
||||
1 - ema_decay))
|
||||
|
||||
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
grad_norm = torch.nn.utils.clip_grad_norm_(transformer.parameters(),
|
||||
max_norm=max_grad_norm)
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
|
||||
return total_loss, grad_norm.item(), model_pred_norm
|
||||
|
||||
|
||||
def main(args):
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
|
||||
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))
|
||||
torch.cuda.set_device(rank)
|
||||
init_distributed_environment(world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank)
|
||||
initialize_model_parallel(tensor_model_parallel_size=args.sp_size,
|
||||
sequence_model_parallel_size=args.sp_size)
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=MODEL_PATH,
|
||||
num_gpus=world_size,
|
||||
use_cpu_offload=False,
|
||||
precision=args.master_weight_type,
|
||||
dit_config=WanVideoConfig(),
|
||||
device_str="cuda",
|
||||
)
|
||||
fastvideo_args.check_fastvideo_args()
|
||||
|
||||
device_str = f"cuda:{rank}"
|
||||
device = torch.device(device_str)
|
||||
fastvideo_args.device = device
|
||||
|
||||
# If passed along, set the training seed now. On GPU...
|
||||
if args.seed is not None:
|
||||
# TODO: t within the same seq parallel group should be the same. Noise should be different.
|
||||
set_seed(args.seed + rank)
|
||||
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
# as these weights are only used for inference, keeping weights in full precision is not required.
|
||||
|
||||
# Create model:
|
||||
|
||||
logger.info("--> loading model from %s", TRANSFORMER_PATH)
|
||||
|
||||
fastvideo_args.device = device
|
||||
transformer_loader = TransformerLoader()
|
||||
transformer = transformer_loader.load(TRANSFORMER_PATH, "", fastvideo_args)
|
||||
transformer = transformer.train()
|
||||
transformer.requires_grad_(True)
|
||||
|
||||
teacher_loader = TransformerLoader()
|
||||
teacher_transformer = teacher_loader.load(TRANSFORMER_PATH, "",
|
||||
fastvideo_args)
|
||||
if args.use_ema:
|
||||
ema_transformer = teacher_loader.load(TRANSFORMER_PATH, "",
|
||||
fastvideo_args)
|
||||
else:
|
||||
ema_transformer = None
|
||||
|
||||
logger.info(
|
||||
" Total training parameters = %s M",
|
||||
sum(p.numel()
|
||||
for p in transformer.parameters() if p.requires_grad) / 1e6)
|
||||
logger.info("--> model loaded")
|
||||
|
||||
teacher_transformer.requires_grad_(False)
|
||||
if args.use_ema:
|
||||
ema_transformer.requires_grad_(False)
|
||||
|
||||
# scheduler
|
||||
noise_scheduler_loader = SchedulerLoader()
|
||||
noise_scheduler = noise_scheduler_loader.load(SCHEDULER_PATH, "",
|
||||
fastvideo_args)
|
||||
solver = EulerSolver(
|
||||
noise_scheduler.sigmas.numpy()[::-1],
|
||||
noise_scheduler.config.num_train_timesteps,
|
||||
euler_timesteps=args.num_euler_timesteps,
|
||||
)
|
||||
solver.to(device)
|
||||
params_to_optimize = transformer.parameters()
|
||||
params_to_optimize = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
|
||||
optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
lr=args.learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
init_steps = 0
|
||||
logger.info("optimizer: %s", optimizer)
|
||||
|
||||
# todo add lr scheduler
|
||||
lr_scheduler = get_scheduler(
|
||||
args.lr_scheduler,
|
||||
optimizer=optimizer,
|
||||
num_warmup_steps=args.lr_warmup_steps * world_size,
|
||||
num_training_steps=args.max_train_steps * world_size,
|
||||
num_cycles=args.lr_num_cycles,
|
||||
power=args.lr_power,
|
||||
last_epoch=init_steps - 1,
|
||||
)
|
||||
|
||||
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
|
||||
args.cfg)
|
||||
uncond_prompt_embed = train_dataset.uncond_prompt_embed
|
||||
uncond_prompt_mask = train_dataset.uncond_prompt_mask
|
||||
sampler = (LengthGroupedSampler(
|
||||
args.train_batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
lengths=train_dataset.lengths,
|
||||
group_frame=args.group_frame,
|
||||
group_resolution=args.group_resolution,
|
||||
) if (args.group_frame or args.group_resolution) else DistributedSampler(
|
||||
train_dataset, rank=rank, num_replicas=world_size, shuffle=False))
|
||||
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
collate_fn=latent_collate_function,
|
||||
pin_memory=True,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
drop_last=True,
|
||||
)
|
||||
|
||||
num_update_steps_per_epoch = math.ceil(
|
||||
len(train_dataloader) / args.gradient_accumulation_steps *
|
||||
args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps /
|
||||
num_update_steps_per_epoch)
|
||||
|
||||
if rank <= 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
# Train!
|
||||
total_batch_size = (world_size * args.gradient_accumulation_steps /
|
||||
args.sp_size * args.train_sp_batch_size)
|
||||
logger.info("***** Running training *****")
|
||||
logger.info(" Num examples = %s", len(train_dataset))
|
||||
logger.info(" Dataloader size = %s", len(train_dataloader))
|
||||
logger.info(" Num Epochs = %s", args.num_train_epochs)
|
||||
logger.info(" Resume training from step %s", init_steps)
|
||||
logger.info(" Instantaneous batch size per device = %s",
|
||||
args.train_batch_size)
|
||||
logger.info(
|
||||
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
|
||||
total_batch_size)
|
||||
logger.info(" Gradient Accumulation steps = %s",
|
||||
args.gradient_accumulation_steps)
|
||||
logger.info(" Total optimization steps = %s", args.max_train_steps)
|
||||
logger.info(
|
||||
" Total training parameters per FSDP shard = %s B",
|
||||
sum(p.numel()
|
||||
for p in transformer.parameters() if p.requires_grad) / 1e9)
|
||||
# print dtype
|
||||
logger.info(" Master weight dtype: %s",
|
||||
transformer.parameters().__next__().dtype)
|
||||
|
||||
# Potentially load in the weights and states from a previous save
|
||||
if args.resume_from_checkpoint:
|
||||
assert NotImplementedError(
|
||||
"resume_from_checkpoint is not supported now.")
|
||||
# TODO
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, args.max_train_steps),
|
||||
initial=init_steps,
|
||||
desc="Steps",
|
||||
# Only show the progress bar once on each machine.
|
||||
disable=local_rank > 0,
|
||||
)
|
||||
|
||||
loader = sp_parallel_dataloader_wrapper(
|
||||
train_dataloader,
|
||||
device,
|
||||
args.train_batch_size,
|
||||
args.sp_size,
|
||||
args.train_sp_batch_size,
|
||||
)
|
||||
|
||||
step_times = deque(maxlen=100)
|
||||
|
||||
# todo future
|
||||
for i in range(init_steps):
|
||||
next(loader)
|
||||
|
||||
# log_validation(args, transformer, device,
|
||||
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
|
||||
def get_num_phases(multi_phased_distill_schedule, step):
|
||||
# step-phase,step-phase
|
||||
multi_phases = multi_phased_distill_schedule.split(",")
|
||||
phase = multi_phases[-1].split("-")[-1]
|
||||
for step_phases in multi_phases:
|
||||
phase_step, phase = step_phases.split("-")
|
||||
if step <= int(phase_step):
|
||||
return int(phase)
|
||||
return phase
|
||||
|
||||
for step in range(init_steps + 1, args.max_train_steps + 1):
|
||||
start_time = time.time()
|
||||
assert args.multi_phased_distill_schedule is not None
|
||||
num_phases = get_num_phases(args.multi_phased_distill_schedule, step)
|
||||
|
||||
loss, grad_norm, pred_norm = distill_one_step(
|
||||
transformer,
|
||||
args.model_type,
|
||||
teacher_transformer,
|
||||
ema_transformer,
|
||||
optimizer,
|
||||
lr_scheduler,
|
||||
loader,
|
||||
noise_scheduler,
|
||||
solver,
|
||||
noise_random_generator,
|
||||
args.gradient_accumulation_steps,
|
||||
args.sp_size,
|
||||
args.max_grad_norm,
|
||||
uncond_prompt_embed,
|
||||
uncond_prompt_mask,
|
||||
args.num_euler_timesteps,
|
||||
num_phases,
|
||||
args.not_apply_cfg_solver,
|
||||
args.distill_cfg,
|
||||
args.ema_decay,
|
||||
args.pred_decay_weight,
|
||||
args.pred_decay_type,
|
||||
args.hunyuan_teacher_disable_cfg,
|
||||
)
|
||||
|
||||
step_time = time.time() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"loss": f"{loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
"phases": num_phases,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss":
|
||||
loss,
|
||||
"learning_rate":
|
||||
lr_scheduler.get_last_lr()[0],
|
||||
"step_time":
|
||||
step_time,
|
||||
"avg_step_time":
|
||||
avg_step_time,
|
||||
"grad_norm":
|
||||
grad_norm,
|
||||
"pred_fro_norm":
|
||||
pred_norm["fro"], # codespell:ignore
|
||||
"pred_largest_singular_value":
|
||||
pred_norm["largest singular value"],
|
||||
"pred_absolute_mean":
|
||||
pred_norm["absolute mean"],
|
||||
"pred_absolute_max":
|
||||
pred_norm["absolute max"],
|
||||
},
|
||||
step=step,
|
||||
)
|
||||
if step % args.checkpointing_steps == 0:
|
||||
if args.use_lora:
|
||||
# Save LoRA weights
|
||||
save_lora_checkpoint(transformer, optimizer, rank,
|
||||
args.output_dir, step)
|
||||
else:
|
||||
# Your existing checkpoint saving code
|
||||
if args.use_ema:
|
||||
save_checkpoint(ema_transformer, rank, args.output_dir,
|
||||
step)
|
||||
else:
|
||||
save_checkpoint(transformer, rank, args.output_dir, step)
|
||||
dist.barrier()
|
||||
if args.log_validation and step % args.validation_steps == 0:
|
||||
log_validation(
|
||||
args,
|
||||
transformer,
|
||||
device,
|
||||
torch.bfloat16,
|
||||
step,
|
||||
scheduler_type=args.scheduler_type,
|
||||
shift=args.shift,
|
||||
num_euler_timesteps=args.num_euler_timesteps,
|
||||
linear_quadratic_threshold=args.linear_quadratic_threshold,
|
||||
linear_range=args.linear_range,
|
||||
ema=False,
|
||||
)
|
||||
if args.use_ema:
|
||||
log_validation(
|
||||
args,
|
||||
ema_transformer,
|
||||
device,
|
||||
torch.bfloat16,
|
||||
step,
|
||||
scheduler_type=args.scheduler_type,
|
||||
shift=args.shift,
|
||||
num_euler_timesteps=args.num_euler_timesteps,
|
||||
linear_quadratic_threshold=args.linear_quadratic_threshold,
|
||||
linear_range=args.linear_range,
|
||||
ema=True,
|
||||
)
|
||||
|
||||
if args.use_lora:
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
|
||||
args.max_train_steps)
|
||||
else:
|
||||
save_checkpoint(transformer, rank, args.output_dir,
|
||||
args.max_train_steps)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
destroy_sequence_parallel_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument("--model_type",
|
||||
type=str,
|
||||
default="mochi",
|
||||
help="The type of model to train.")
|
||||
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--data_json_path", type=str, required=True)
|
||||
parser.add_argument("--num_height", type=int, default=480)
|
||||
parser.add_argument("--num_width", type=int, default=848)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument(
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=10,
|
||||
help=
|
||||
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_batch_size",
|
||||
type=int,
|
||||
default=16,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument("--num_latent_t",
|
||||
type=int,
|
||||
default=28,
|
||||
help="Number of latent timesteps.")
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--pretrained_model_name_or_path", type=str)
|
||||
parser.add_argument("--dit_model_name_or_path", type=str)
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
|
||||
# diffusion setting
|
||||
parser.add_argument("--ema_decay", type=float, default=0.95)
|
||||
parser.add_argument("--ema_start_step", type=int, default=0)
|
||||
parser.add_argument("--cfg", type=float, default=0.1)
|
||||
|
||||
# validation & logs
|
||||
parser.add_argument("--validation_prompt_dir", type=str)
|
||||
parser.add_argument("--validation_sampling_steps", type=str, default="64")
|
||||
parser.add_argument("--validation_guidance_scale", type=str, default="4.5")
|
||||
|
||||
parser.add_argument("--validation_steps", type=float, default=64)
|
||||
parser.add_argument("--log_validation", action="store_true")
|
||||
parser.add_argument("--tracker_project_name", type=str, default=None)
|
||||
parser.add_argument("--seed",
|
||||
type=int,
|
||||
default=None,
|
||||
help="A seed for reproducible training.")
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help=
|
||||
"The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoints_total_limit",
|
||||
type=int,
|
||||
default=None,
|
||||
help=("Max number of checkpoints to store."),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpointing_steps",
|
||||
type=int,
|
||||
default=500,
|
||||
help=
|
||||
("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
|
||||
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
|
||||
" training using `--resume_from_checkpoint`."),
|
||||
)
|
||||
parser.add_argument("--shift", type=float, default=1.0)
|
||||
parser.add_argument(
|
||||
"--resume_from_checkpoint",
|
||||
type=str,
|
||||
default=None,
|
||||
help=
|
||||
("Whether training should be resumed from a previous checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--resume_from_lora_checkpoint",
|
||||
type=str,
|
||||
default=None,
|
||||
help=
|
||||
("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=
|
||||
("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
parser.add_argument(
|
||||
"--max_train_steps",
|
||||
type=int,
|
||||
default=None,
|
||||
help=
|
||||
"Total number of training steps to perform. If provided, overrides num_train_epochs.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--gradient_accumulation_steps",
|
||||
type=int,
|
||||
default=1,
|
||||
help=
|
||||
"Number of updates steps to accumulate before performing a backward/update pass.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--learning_rate",
|
||||
type=float,
|
||||
default=1e-4,
|
||||
help="Initial learning rate (after the potential warmup period) to use.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--scale_lr",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help=
|
||||
"Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lr_warmup_steps",
|
||||
type=int,
|
||||
default=10,
|
||||
help="Number of steps for the warmup in the lr scheduler.",
|
||||
)
|
||||
parser.add_argument("--max_grad_norm",
|
||||
default=1.0,
|
||||
type=float,
|
||||
help="Max gradient norm.")
|
||||
parser.add_argument(
|
||||
"--gradient_checkpointing",
|
||||
action="store_true",
|
||||
help=
|
||||
"Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
|
||||
)
|
||||
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
|
||||
parser.add_argument(
|
||||
"--allow_tf32",
|
||||
action="store_true",
|
||||
help=
|
||||
("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
|
||||
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mixed_precision",
|
||||
type=str,
|
||||
default=None,
|
||||
choices=["no", "fp16", "bf16"],
|
||||
help=
|
||||
("Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
|
||||
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
|
||||
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_cpu_offload",
|
||||
action="store_true",
|
||||
help=
|
||||
"Whether to use CPU offload for param & gradient & optimizer states.",
|
||||
)
|
||||
|
||||
parser.add_argument("--sp_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="For sequence parallel")
|
||||
parser.add_argument(
|
||||
"--train_sp_batch_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Batch size for sequence parallel training",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--use_lora",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Whether to use LoRA for finetuning.",
|
||||
)
|
||||
parser.add_argument("--lora_alpha",
|
||||
type=int,
|
||||
default=256,
|
||||
help="Alpha parameter for LoRA.")
|
||||
parser.add_argument("--lora_rank",
|
||||
type=int,
|
||||
default=128,
|
||||
help="LoRA rank parameter. ")
|
||||
parser.add_argument("--fsdp_sharding_startegy", default="full")
|
||||
|
||||
# lr_scheduler
|
||||
parser.add_argument(
|
||||
"--lr_scheduler",
|
||||
type=str,
|
||||
default="constant",
|
||||
help=
|
||||
('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
|
||||
' "constant", "constant_with_warmup"]'),
|
||||
)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=100)
|
||||
parser.add_argument(
|
||||
"--lr_num_cycles",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of cycles in the learning rate scheduler.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lr_power",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="Power factor of the polynomial scheduler.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--not_apply_cfg_solver",
|
||||
action="store_true",
|
||||
help="Whether to apply the cfg_solver.",
|
||||
)
|
||||
parser.add_argument("--distill_cfg",
|
||||
type=float,
|
||||
default=3.0,
|
||||
help="Distillation coefficient.")
|
||||
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
|
||||
parser.add_argument("--scheduler_type",
|
||||
type=str,
|
||||
default="pcm",
|
||||
help="The scheduler type to use.")
|
||||
parser.add_argument(
|
||||
"--linear_quadratic_threshold",
|
||||
type=float,
|
||||
default=0.025,
|
||||
help="Threshold for linear quadratic scheduler.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--linear_range",
|
||||
type=float,
|
||||
default=0.5,
|
||||
help="Range for linear quadratic scheduler.",
|
||||
)
|
||||
parser.add_argument("--weight_decay",
|
||||
type=float,
|
||||
default=0.001,
|
||||
help="Weight decay to apply.")
|
||||
parser.add_argument("--use_ema",
|
||||
action="store_true",
|
||||
help="Whether to use EMA.")
|
||||
parser.add_argument("--multi_phased_distill_schedule",
|
||||
type=str,
|
||||
default=None)
|
||||
parser.add_argument("--pred_decay_weight", type=float, default=0.0)
|
||||
parser.add_argument("--pred_decay_type", default="l1")
|
||||
parser.add_argument("--hunyuan_teacher_disable_cfg", action="store_true")
|
||||
parser.add_argument(
|
||||
"--master_weight_type",
|
||||
type=str,
|
||||
default="fp32",
|
||||
help="Weight type to use - fp32 or bf16.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -237,7 +237,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
|
||||
type=str,
|
||||
default="540p",
|
||||
choices=["540p", "720p"],
|
||||
help="The resolution of the model.",
|
||||
help="Root path of all the models, including t2v models and extra models.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--load-key",
|
||||
@@ -361,7 +361,7 @@ def add_parallel_args(parser: argparse.ArgumentParser):
|
||||
"--ring-degree",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Ring degree.",
|
||||
help="Ulysses degree.",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
@@ -17,7 +17,7 @@ from fastvideo.models.hunyuan.vae import load_vae
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
|
||||
|
||||
class Inference:
|
||||
class Inference(object):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
||||
@@ -41,7 +41,7 @@ def get_rewrite_prompt(ori_prompt, mode="Normal"):
|
||||
elif mode == "Master":
|
||||
prompt = master_mode_prompt.format(input=ori_prompt)
|
||||
else:
|
||||
raise Exception("Only supports Normal and Master mode, but got {}".format(mode))
|
||||
raise Exception("Only supports Normal and Normal", mode)
|
||||
return prompt
|
||||
|
||||
|
||||
|
||||
@@ -31,7 +31,7 @@ mochi_latents_std = torch.tensor([
|
||||
mochi_scaling_factor = 1.0
|
||||
|
||||
|
||||
def normalize_dit_input(model_type, latents):
|
||||
def normalize_dit_input(model_type, latents, args=None):
|
||||
if model_type == "mochi":
|
||||
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
|
||||
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
|
||||
@@ -41,5 +41,16 @@ def normalize_dit_input(model_type, latents):
|
||||
return latents * 0.476986
|
||||
elif model_type == "hunyuan":
|
||||
return latents * 0.476986
|
||||
elif model_type == "wan":
|
||||
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
|
||||
vae_config = WanVAEConfig()
|
||||
latents_mean = torch.tensor(vae_config.arch_config.latents_mean)
|
||||
latents_std = 1.0 / torch.tensor(vae_config.arch_config.latents_std)
|
||||
|
||||
|
||||
latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(device=latents.device)
|
||||
latents_std = latents_std.view(1, -1, 1, 1, 1).to(device=latents.device)
|
||||
latents = ((latents.float() - latents_mean) * latents_std).to(latents)
|
||||
return latents
|
||||
else:
|
||||
raise NotImplementedError(f"model_type {model_type} not supported")
|
||||
|
||||
@@ -267,25 +267,25 @@ class Step1Model(PreTrainedModel):
|
||||
class STEP1TextEncoder(torch.nn.Module):
|
||||
|
||||
def __init__(self, model_dir, max_length=320):
|
||||
super()
|
||||
super(STEP1TextEncoder, self).__init__()
|
||||
self.max_length = max_length
|
||||
self.text_tokenizer = Wrapped_StepChatTokenizer(os.path.join(model_dir, 'step1_chat_tokenizer.model'))
|
||||
text_encoder = Step1Model.from_pretrained(model_dir)
|
||||
self.text_encoder = text_encoder.eval().to(torch.bfloat16)
|
||||
|
||||
@torch.no_grad
|
||||
@torch.autocast(device_type='cuda', dtype=torch.bfloat16)
|
||||
def forward(self, prompts, with_mask=True, max_length=None):
|
||||
self.device = next(self.text_encoder.parameters()).device
|
||||
if type(prompts) is str:
|
||||
prompts = [prompts]
|
||||
with torch.no_grad(), torch.cuda.amp.autocast(dtype=torch.bfloat16):
|
||||
if type(prompts) is str:
|
||||
prompts = [prompts]
|
||||
|
||||
txt_tokens = self.text_tokenizer(prompts,
|
||||
max_length=max_length or self.max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_tensors="pt")
|
||||
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
|
||||
txt_tokens = self.text_tokenizer(prompts,
|
||||
max_length=max_length or self.max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_tensors="pt")
|
||||
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
|
||||
attention_mask=txt_tokens.attention_mask.to(self.device) if with_mask else None)
|
||||
y_mask = txt_tokens.attention_mask
|
||||
y_mask = txt_tokens.attention_mask
|
||||
return y.transpose(0, 1), y_mask
|
||||
|
||||
@@ -1,486 +0,0 @@
|
||||
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import math
|
||||
from typing import Any, Dict, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
|
||||
from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers
|
||||
from diffusers.models.attention import FeedForward
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.models.cache_utils import CacheMixin
|
||||
from diffusers.models.embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps, get_1d_rotary_pos_embed
|
||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.normalization import FP32LayerNorm
|
||||
|
||||
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
class WanAttnProcessor2_0:
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("WanAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
encoder_hidden_states_img = None
|
||||
if attn.add_k_proj is not None:
|
||||
# 512 is the context length of the text encoder, hardcoded for now
|
||||
image_context_length = encoder_hidden_states.shape[1] - 512
|
||||
encoder_hidden_states_img = encoder_hidden_states[:, :image_context_length]
|
||||
encoder_hidden_states = encoder_hidden_states[:, image_context_length:]
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
|
||||
if rotary_emb is not None:
|
||||
|
||||
def apply_rotary_emb(hidden_states: torch.Tensor, freqs: torch.Tensor):
|
||||
x_rotated = torch.view_as_complex(hidden_states.to(torch.float64).unflatten(3, (-1, 2)))
|
||||
x_out = torch.view_as_real(x_rotated * freqs).flatten(3, 4)
|
||||
return x_out.type_as(hidden_states)
|
||||
|
||||
query = apply_rotary_emb(query, rotary_emb)
|
||||
key = apply_rotary_emb(key, rotary_emb)
|
||||
|
||||
# I2V task
|
||||
hidden_states_img = None
|
||||
if encoder_hidden_states_img is not None:
|
||||
key_img = attn.add_k_proj(encoder_hidden_states_img)
|
||||
key_img = attn.norm_added_k(key_img)
|
||||
value_img = attn.add_v_proj(encoder_hidden_states_img)
|
||||
|
||||
key_img = key_img.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
value_img = value_img.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
|
||||
hidden_states_img = F.scaled_dot_product_attention(
|
||||
query, key_img, value_img, attn_mask=None, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
hidden_states_img = hidden_states_img.transpose(1, 2).flatten(2, 3)
|
||||
hidden_states_img = hidden_states_img.type_as(query)
|
||||
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
hidden_states = hidden_states.transpose(1, 2).flatten(2, 3)
|
||||
hidden_states = hidden_states.type_as(query)
|
||||
|
||||
if hidden_states_img is not None:
|
||||
hidden_states = hidden_states + hidden_states_img
|
||||
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanImageEmbedding(torch.nn.Module):
|
||||
def __init__(self, in_features: int, out_features: int, pos_embed_seq_len=None):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = FP32LayerNorm(in_features)
|
||||
self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu")
|
||||
self.norm2 = FP32LayerNorm(out_features)
|
||||
if pos_embed_seq_len is not None:
|
||||
self.pos_embed = nn.Parameter(torch.zeros(1, pos_embed_seq_len, in_features))
|
||||
else:
|
||||
self.pos_embed = None
|
||||
|
||||
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
if self.pos_embed is not None:
|
||||
batch_size, seq_len, embed_dim = encoder_hidden_states_image.shape
|
||||
encoder_hidden_states_image = encoder_hidden_states_image.view(-1, 2 * seq_len, embed_dim)
|
||||
encoder_hidden_states_image = encoder_hidden_states_image + self.pos_embed
|
||||
|
||||
hidden_states = self.norm1(encoder_hidden_states_image)
|
||||
hidden_states = self.ff(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanTimeTextImageEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
time_freq_dim: int,
|
||||
time_proj_dim: int,
|
||||
text_embed_dim: int,
|
||||
image_embed_dim: Optional[int] = None,
|
||||
pos_embed_seq_len: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0)
|
||||
self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim)
|
||||
self.act_fn = nn.SiLU()
|
||||
self.time_proj = nn.Linear(dim, time_proj_dim)
|
||||
self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh")
|
||||
|
||||
self.image_embedder = None
|
||||
if image_embed_dim is not None:
|
||||
self.image_embedder = WanImageEmbedding(image_embed_dim, dim, pos_embed_seq_len=pos_embed_seq_len)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_image: Optional[torch.Tensor] = None,
|
||||
):
|
||||
timestep = self.timesteps_proj(timestep)
|
||||
|
||||
time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype
|
||||
if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8:
|
||||
timestep = timestep.to(time_embedder_dtype)
|
||||
temb = self.time_embedder(timestep).type_as(encoder_hidden_states)
|
||||
timestep_proj = self.time_proj(self.act_fn(temb))
|
||||
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image)
|
||||
|
||||
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
|
||||
|
||||
|
||||
class WanRotaryPosEmbed(nn.Module):
|
||||
def __init__(
|
||||
self, attention_head_dim: int, patch_size: Tuple[int, int, int], max_seq_len: int, theta: float = 10000.0
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.attention_head_dim = attention_head_dim
|
||||
self.patch_size = patch_size
|
||||
self.max_seq_len = max_seq_len
|
||||
|
||||
h_dim = w_dim = 2 * (attention_head_dim // 6)
|
||||
t_dim = attention_head_dim - h_dim - w_dim
|
||||
|
||||
freqs = []
|
||||
for dim in [t_dim, h_dim, w_dim]:
|
||||
freq = get_1d_rotary_pos_embed(
|
||||
dim, max_seq_len, theta, use_real=False, repeat_interleave_real=False, freqs_dtype=torch.float64
|
||||
)
|
||||
freqs.append(freq)
|
||||
self.freqs = torch.cat(freqs, dim=1)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
ppf, pph, ppw = num_frames // p_t, height // p_h, width // p_w
|
||||
|
||||
freqs = self.freqs.to(hidden_states.device)
|
||||
freqs = freqs.split_with_sizes(
|
||||
[
|
||||
self.attention_head_dim // 2 - 2 * (self.attention_head_dim // 6),
|
||||
self.attention_head_dim // 6,
|
||||
self.attention_head_dim // 6,
|
||||
],
|
||||
dim=1,
|
||||
)
|
||||
|
||||
freqs_f = freqs[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1)
|
||||
freqs_h = freqs[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1)
|
||||
freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1)
|
||||
freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1)
|
||||
return freqs
|
||||
|
||||
|
||||
class WanTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.attn1 = Attention(
|
||||
query_dim=dim,
|
||||
heads=num_heads,
|
||||
kv_heads=num_heads,
|
||||
dim_head=dim // num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps,
|
||||
bias=True,
|
||||
cross_attention_dim=None,
|
||||
out_bias=True,
|
||||
processor=WanAttnProcessor2_0(),
|
||||
)
|
||||
|
||||
# 2. Cross-attention
|
||||
self.attn2 = Attention(
|
||||
query_dim=dim,
|
||||
heads=num_heads,
|
||||
kv_heads=num_heads,
|
||||
dim_head=dim // num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps,
|
||||
bias=True,
|
||||
cross_attention_dim=None,
|
||||
out_bias=True,
|
||||
added_kv_proj_dim=added_kv_proj_dim,
|
||||
added_proj_bias=True,
|
||||
processor=WanAttnProcessor2_0(),
|
||||
)
|
||||
self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate")
|
||||
self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
rotary_emb: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
|
||||
self.scale_shift_table + temb.float()
|
||||
).chunk(6, dim=1)
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states)
|
||||
attn_output = self.attn1(hidden_states=norm_hidden_states, rotary_emb=rotary_emb)
|
||||
hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states)
|
||||
|
||||
# 2. Cross-attention
|
||||
norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states)
|
||||
attn_output = self.attn2(hidden_states=norm_hidden_states, encoder_hidden_states=encoder_hidden_states)
|
||||
hidden_states = hidden_states + attn_output
|
||||
|
||||
# 3. Feed-forward
|
||||
norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as(
|
||||
hidden_states
|
||||
)
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin):
|
||||
r"""
|
||||
A Transformer model for video-like data used in the Wan model.
|
||||
|
||||
Args:
|
||||
patch_size (`Tuple[int]`, defaults to `(1, 2, 2)`):
|
||||
3D patch dimensions for video embedding (t_patch, h_patch, w_patch).
|
||||
num_attention_heads (`int`, defaults to `40`):
|
||||
Fixed length for text embeddings.
|
||||
attention_head_dim (`int`, defaults to `128`):
|
||||
The number of channels in each head.
|
||||
in_channels (`int`, defaults to `16`):
|
||||
The number of channels in the input.
|
||||
out_channels (`int`, defaults to `16`):
|
||||
The number of channels in the output.
|
||||
text_dim (`int`, defaults to `512`):
|
||||
Input dimension for text embeddings.
|
||||
freq_dim (`int`, defaults to `256`):
|
||||
Dimension for sinusoidal time embeddings.
|
||||
ffn_dim (`int`, defaults to `13824`):
|
||||
Intermediate dimension in feed-forward network.
|
||||
num_layers (`int`, defaults to `40`):
|
||||
The number of layers of transformer blocks to use.
|
||||
window_size (`Tuple[int]`, defaults to `(-1, -1)`):
|
||||
Window size for local attention (-1 indicates global attention).
|
||||
cross_attn_norm (`bool`, defaults to `True`):
|
||||
Enable cross-attention normalization.
|
||||
qk_norm (`bool`, defaults to `True`):
|
||||
Enable query/key normalization.
|
||||
eps (`float`, defaults to `1e-6`):
|
||||
Epsilon value for normalization layers.
|
||||
add_img_emb (`bool`, defaults to `False`):
|
||||
Whether to use img_emb.
|
||||
added_kv_proj_dim (`int`, *optional*, defaults to `None`):
|
||||
The number of channels to use for the added key and value projections. If `None`, no projection is used.
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
_skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"]
|
||||
_no_split_modules = ["WanTransformerBlock"]
|
||||
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
|
||||
_keys_to_ignore_on_load_unexpected = ["norm_added_q"]
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: Tuple[int] = (1, 2, 2),
|
||||
num_attention_heads: int = 40,
|
||||
attention_head_dim: int = 128,
|
||||
in_channels: int = 16,
|
||||
out_channels: int = 16,
|
||||
text_dim: int = 4096,
|
||||
freq_dim: int = 256,
|
||||
ffn_dim: int = 13824,
|
||||
num_layers: int = 40,
|
||||
cross_attn_norm: bool = True,
|
||||
qk_norm: Optional[str] = "rms_norm_across_heads",
|
||||
eps: float = 1e-6,
|
||||
image_dim: Optional[int] = None,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
rope_max_seq_len: int = 1024,
|
||||
pos_embed_seq_len: Optional[int] = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
out_channels = out_channels or in_channels
|
||||
|
||||
# 1. Patch & position embedding
|
||||
self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len)
|
||||
self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size)
|
||||
|
||||
# 2. Condition embeddings
|
||||
# image_embedding_dim=1280 for I2V model
|
||||
self.condition_embedder = WanTimeTextImageEmbedding(
|
||||
dim=inner_dim,
|
||||
time_freq_dim=freq_dim,
|
||||
time_proj_dim=inner_dim * 6,
|
||||
text_embed_dim=text_dim,
|
||||
image_embed_dim=image_dim,
|
||||
pos_embed_seq_len=pos_embed_seq_len,
|
||||
)
|
||||
|
||||
# 3. Transformer blocks
|
||||
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)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Output norm & projection
|
||||
self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False)
|
||||
self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size))
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_image: Optional[torch.Tensor] = None,
|
||||
return_dict: bool = True,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
if attention_kwargs is not None:
|
||||
attention_kwargs = attention_kwargs.copy()
|
||||
lora_scale = attention_kwargs.pop("scale", 1.0)
|
||||
else:
|
||||
lora_scale = 1.0
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# weight the lora layers by setting `lora_scale` for each PEFT layer
|
||||
scale_lora_layers(self, lora_scale)
|
||||
else:
|
||||
if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None:
|
||||
logger.warning(
|
||||
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
|
||||
)
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.config.patch_size
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
rotary_emb = self.rope(hidden_states)
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image
|
||||
)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block, hidden_states, encoder_hidden_states, timestep_proj, rotary_emb
|
||||
)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb)
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
|
||||
|
||||
# Move the shift and scale tensors to the same device as hidden_states.
|
||||
# When using multi-GPU inference via accelerate these will be on the
|
||||
# first device rather than the last device, which hidden_states ends up
|
||||
# on.
|
||||
shift = shift.to(hidden_states.device)
|
||||
scale = scale.to(hidden_states.device)
|
||||
|
||||
hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(
|
||||
batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1
|
||||
)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# remove `lora_scale` from each PEFT layer
|
||||
unscale_lora_layers(self, lora_scale)
|
||||
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
|
||||
return Transformer2DModelOutput(sample=output)
|
||||
@@ -1,609 +0,0 @@
|
||||
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import html
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
|
||||
import regex as re
|
||||
import torch
|
||||
from transformers import AutoTokenizer, UMT5EncoderModel
|
||||
|
||||
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
||||
from diffusers.loaders import WanLoraLoaderMixin
|
||||
from diffusers.models import AutoencoderKLWan, WanTransformer3DModel
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import is_ftfy_available, is_torch_xla_available, logging, replace_example_docstring
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.pipelines.wan.pipeline_output import WanPipelineOutput
|
||||
from einops import rearrange
|
||||
from transformers import UMT5EncoderModel, T5TokenizerFast
|
||||
|
||||
from fastvideo.models.mochi_hf.modeling_wan import WanTransformer3DModel
|
||||
from fastvideo.utils.communications import all_gather
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
|
||||
XLA_AVAILABLE = True
|
||||
else:
|
||||
XLA_AVAILABLE = False
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
if is_ftfy_available():
|
||||
import ftfy
|
||||
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
Examples:
|
||||
```python
|
||||
>>> import torch
|
||||
>>> from diffusers.utils import export_to_video
|
||||
>>> from diffusers import AutoencoderKLWan, WanPipeline
|
||||
>>> from diffusers.schedulers.scheduling_unipc_multistep import UniPCMultistepScheduler
|
||||
|
||||
>>> # Available models: Wan-AI/Wan2.1-T2V-14B-Diffusers, Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
>>> model_id = "Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
>>> vae = AutoencoderKLWan.from_pretrained(model_id, subfolder="vae", torch_dtype=torch.float32)
|
||||
>>> pipe = WanPipeline.from_pretrained(model_id, vae=vae, torch_dtype=torch.bfloat16)
|
||||
>>> flow_shift = 5.0 # 5.0 for 720P, 3.0 for 480P
|
||||
>>> pipe.scheduler = UniPCMultistepScheduler.from_config(pipe.scheduler.config, flow_shift=flow_shift)
|
||||
>>> pipe.to("cuda")
|
||||
|
||||
>>> prompt = "A cat and a dog baking a cake together in a kitchen. The cat is carefully measuring flour, while the dog is stirring the batter with a wooden spoon. The kitchen is cozy, with sunlight streaming through the window."
|
||||
>>> negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
|
||||
>>> output = pipe(
|
||||
... prompt=prompt,
|
||||
... negative_prompt=negative_prompt,
|
||||
... height=720,
|
||||
... width=1280,
|
||||
... num_frames=81,
|
||||
... guidance_scale=5.0,
|
||||
... ).frames[0]
|
||||
>>> export_to_video(output, "output.mp4", fps=16)
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
def basic_clean(text):
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
|
||||
def whitespace_clean(text):
|
||||
text = re.sub(r"\s+", " ", text)
|
||||
text = text.strip()
|
||||
return text
|
||||
|
||||
|
||||
def prompt_clean(text):
|
||||
text = whitespace_clean(basic_clean(text))
|
||||
return text
|
||||
|
||||
|
||||
class WanPipeline(DiffusionPipeline, WanLoraLoaderMixin):
|
||||
r"""
|
||||
Pipeline for text-to-video generation using Wan.
|
||||
|
||||
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
|
||||
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
|
||||
|
||||
Args:
|
||||
tokenizer ([`T5Tokenizer`]):
|
||||
Tokenizer from [T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5Tokenizer),
|
||||
specifically the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant.
|
||||
text_encoder ([`T5EncoderModel`]):
|
||||
[T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
|
||||
the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant.
|
||||
transformer ([`WanTransformer3DModel`]):
|
||||
Conditional Transformer to denoise the input latents.
|
||||
scheduler ([`UniPCMultistepScheduler`]):
|
||||
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
|
||||
vae ([`AutoencoderKLWan`]):
|
||||
Variational Auto-Encoder (VAE) Model to encode and decode videos to and from latent representations.
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "text_encoder->transformer->vae"
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: AutoTokenizer,
|
||||
text_encoder: UMT5EncoderModel,
|
||||
transformer: WanTransformer3DModel,
|
||||
vae: AutoencoderKLWan,
|
||||
scheduler: FlowMatchEulerDiscreteScheduler,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
|
||||
self.vae_scale_factor_temporal = 2 ** sum(self.vae.temperal_downsample) if getattr(self, "vae", None) else 4
|
||||
self.vae_scale_factor_spatial = 2 ** len(self.vae.temperal_downsample) if getattr(self, "vae", None) else 8
|
||||
self.video_processor = VideoProcessor(vae_scale_factor=self.vae_scale_factor_spatial)
|
||||
|
||||
def _get_t5_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
num_videos_per_prompt: int = 1,
|
||||
max_sequence_length: int = 226,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
prompt = [prompt_clean(u) for u in prompt]
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_attention_mask=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), mask.to(device)).last_hidden_state
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
|
||||
prompt_embeds = torch.stack(
|
||||
[torch.cat([u, u.new_zeros(max_sequence_length - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0
|
||||
)
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
do_classifier_free_guidance: bool = True,
|
||||
num_videos_per_prompt: int = 1,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
max_sequence_length: int = 226,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
r"""
|
||||
Encodes the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
prompt to be encoded
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts not to guide the image generation. If not defined, one has to pass
|
||||
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
|
||||
less than `1`).
|
||||
do_classifier_free_guidance (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use classifier free guidance or not.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
Number of videos that should be generated per prompt. torch device to place the resulting embeddings on
|
||||
prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
negative_prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
|
||||
argument.
|
||||
device: (`torch.device`, *optional*):
|
||||
torch device
|
||||
dtype: (`torch.dtype`, *optional*):
|
||||
torch dtype
|
||||
"""
|
||||
device = device or self._execution_device
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
if prompt is not None:
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt_embeds = self._get_t5_prompt_embeds(
|
||||
prompt=prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
if do_classifier_free_guidance and negative_prompt_embeds is None:
|
||||
negative_prompt = negative_prompt or ""
|
||||
negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt
|
||||
|
||||
if prompt is not None and type(prompt) is not type(negative_prompt):
|
||||
raise TypeError(
|
||||
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
|
||||
f" {type(prompt)}."
|
||||
)
|
||||
elif batch_size != len(negative_prompt):
|
||||
raise ValueError(
|
||||
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
|
||||
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
|
||||
" the batch size of `prompt`."
|
||||
)
|
||||
|
||||
negative_prompt_embeds = self._get_t5_prompt_embeds(
|
||||
prompt=negative_prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
return prompt_embeds, negative_prompt_embeds
|
||||
|
||||
def check_inputs(
|
||||
self,
|
||||
prompt,
|
||||
negative_prompt,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds=None,
|
||||
negative_prompt_embeds=None,
|
||||
callback_on_step_end_tensor_inputs=None,
|
||||
):
|
||||
if height % 16 != 0 or width % 16 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 16 but are {height} and {width}.")
|
||||
|
||||
if callback_on_step_end_tensor_inputs is not None and not all(
|
||||
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
|
||||
)
|
||||
|
||||
if prompt is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif negative_prompt is not None and negative_prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`: {negative_prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif prompt is None and prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
|
||||
)
|
||||
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
|
||||
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||
elif negative_prompt is not None and (
|
||||
not isinstance(negative_prompt, str) and not isinstance(negative_prompt, list)
|
||||
):
|
||||
raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}")
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size: int,
|
||||
num_channels_latents: int = 16,
|
||||
height: int = 480,
|
||||
width: int = 832,
|
||||
num_frames: int = 81,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if latents is not None:
|
||||
return latents.to(device=device, dtype=dtype)
|
||||
|
||||
num_latent_frames = (num_frames - 1) // self.vae_scale_factor_temporal + 1
|
||||
shape = (
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
num_latent_frames,
|
||||
int(height) // self.vae_scale_factor_spatial,
|
||||
int(width) // self.vae_scale_factor_spatial,
|
||||
)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
return latents
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
@property
|
||||
def do_classifier_free_guidance(self):
|
||||
return self._guidance_scale > 1.0
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
@property
|
||||
def current_timestep(self):
|
||||
return self._current_timestep
|
||||
|
||||
@property
|
||||
def interrupt(self):
|
||||
return self._interrupt
|
||||
|
||||
@property
|
||||
def attention_kwargs(self):
|
||||
return self._attention_kwargs
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
negative_prompt: Union[str, List[str]] = None,
|
||||
height: int = 480,
|
||||
width: int = 832,
|
||||
num_frames: int = 81,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 5.0,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
output_type: Optional[str] = "np",
|
||||
return_dict: bool = True,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
callback_on_step_end: Optional[
|
||||
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
|
||||
] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 512,
|
||||
):
|
||||
r"""
|
||||
The call function to the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
height (`int`, defaults to `480`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, defaults to `832`):
|
||||
The width in pixels of the generated image.
|
||||
num_frames (`int`, defaults to `81`):
|
||||
The number of frames in the generated video.
|
||||
num_inference_steps (`int`, defaults to `50`):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
guidance_scale (`float`, defaults to `5.0`):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion
|
||||
Guidance](https://huggingface.co/papers/2207.12598). `guidance_scale` is defined as `w` of equation 2.
|
||||
of [Imagen Paper](https://huggingface.co/papers/2205.11487). Guidance scale is enabled by setting
|
||||
`guidance_scale > 1`. Higher guidance scale encourages to generate images that are closely linked to
|
||||
the text `prompt`, usually at the expense of lower image quality.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
|
||||
generation deterministic.
|
||||
latents (`torch.Tensor`, *optional*):
|
||||
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor is generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not
|
||||
provided, text embeddings are generated from the `prompt` input argument.
|
||||
output_type (`str`, *optional*, defaults to `"np"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`WanPipelineOutput`] instead of a plain tuple.
|
||||
attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
callback_on_step_end (`Callable`, `PipelineCallback`, `MultiPipelineCallbacks`, *optional*):
|
||||
A function or a subclass of `PipelineCallback` or `MultiPipelineCallbacks` that is called at the end of
|
||||
each denoising step during the inference. with the following arguments: `callback_on_step_end(self:
|
||||
DiffusionPipeline, step: int, timestep: int, callback_kwargs: Dict)`. `callback_kwargs` will include a
|
||||
list of all tensors as specified by `callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
autocast_dtype (`torch.dtype`, *optional*, defaults to `torch.bfloat16`):
|
||||
The dtype to use for the torch.amp.autocast.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~WanPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`, [`WanPipelineOutput`] is returned, otherwise a `tuple` is returned where
|
||||
the first element is a list with the generated images and the second element is a list of `bool`s
|
||||
indicating whether the corresponding generated image contains "not-safe-for-work" (nsfw) content.
|
||||
"""
|
||||
|
||||
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
|
||||
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
callback_on_step_end_tensor_inputs,
|
||||
)
|
||||
|
||||
if num_frames % self.vae_scale_factor_temporal != 1:
|
||||
logger.warning(
|
||||
f"`num_frames - 1` has to be divisible by {self.vae_scale_factor_temporal}. Rounding to the nearest number."
|
||||
)
|
||||
num_frames = num_frames // self.vae_scale_factor_temporal * self.vae_scale_factor_temporal + 1
|
||||
num_frames = max(num_frames, 1)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._attention_kwargs = attention_kwargs
|
||||
self._current_timestep = None
|
||||
self._interrupt = False
|
||||
|
||||
device = self._execution_device
|
||||
|
||||
# 2. Define call parameters
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
# 3. Encode input prompt
|
||||
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
do_classifier_free_guidance=self.do_classifier_free_guidance,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
)
|
||||
|
||||
transformer_dtype = self.transformer.dtype
|
||||
prompt_embeds = prompt_embeds.to(transformer_dtype)
|
||||
if negative_prompt_embeds is not None:
|
||||
negative_prompt_embeds = negative_prompt_embeds.to(transformer_dtype)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
torch.float32,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
|
||||
if get_sequence_parallel_state():
|
||||
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
|
||||
latents = latents[:, :, rank, :, :, :]
|
||||
|
||||
# 6. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
self._progress_bar_config = {"disable": nccl_info.rank_within_group != 0}
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
self._current_timestep = t
|
||||
latent_model_input = latents.to(transformer_dtype)
|
||||
timestep = t.expand(latents.shape[0])
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if self.do_classifier_free_guidance:
|
||||
noise_uncond = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=negative_prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred = noise_uncond + guidance_scale * (noise_pred - noise_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
if XLA_AVAILABLE:
|
||||
xm.mark_step()
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
latents = all_gather(latents, dim=2)
|
||||
|
||||
self._current_timestep = None
|
||||
|
||||
if not output_type == "latent":
|
||||
latents = latents.to(self.vae.dtype)
|
||||
latents_mean = (
|
||||
torch.tensor(self.vae.config.latents_mean)
|
||||
.view(1, self.vae.config.z_dim, 1, 1, 1)
|
||||
.to(latents.device, latents.dtype)
|
||||
)
|
||||
latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to(
|
||||
latents.device, latents.dtype
|
||||
)
|
||||
latents = latents / latents_std + latents_mean
|
||||
video = self.vae.decode(latents, return_dict=False)[0]
|
||||
video = self.video_processor.postprocess_video(video, output_type=output_type)
|
||||
else:
|
||||
video = latents
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (video,)
|
||||
|
||||
return WanPipelineOutput(frames=video)
|
||||
@@ -86,7 +86,7 @@ def inference(args):
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
generator=generator,
|
||||
).frames
|
||||
if nccl_info.global_rank == 0:
|
||||
if nccl_info.global_rank <= 0:
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(
|
||||
@@ -107,7 +107,7 @@ def inference(args):
|
||||
generator=generator,
|
||||
).frames
|
||||
|
||||
if nccl_info.global_rank == 0:
|
||||
if nccl_info.global_rank <= 0:
|
||||
export_to_video(videos[0], args.output_path + ".mp4", fps=24)
|
||||
|
||||
|
||||
|
||||
@@ -94,7 +94,7 @@ def main(args):
|
||||
guidance_scale=args.guidance_scale,
|
||||
generator=generator,
|
||||
).frames
|
||||
if nccl_info.global_rank == 0:
|
||||
if nccl_info.global_rank <= 0:
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(
|
||||
@@ -116,7 +116,7 @@ def main(args):
|
||||
generator=generator,
|
||||
).frames
|
||||
|
||||
if nccl_info.global_rank == 0:
|
||||
if nccl_info.global_rank <= 0:
|
||||
export_to_video(videos[0], args.output_path + ".mp4", fps=30)
|
||||
|
||||
|
||||
|
||||
+11
-5
@@ -20,7 +20,7 @@ from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.utils.latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.models.hunyuan_hf.pipeline_hunyuan import HunyuanVideoPipeline
|
||||
|
||||
@@ -185,7 +185,7 @@ def main(args):
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank == 0 and args.output_dir is not None:
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
@@ -316,7 +316,7 @@ def main(args):
|
||||
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
if rank == 0:
|
||||
if rank <= 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
@@ -393,7 +393,7 @@ def main(args):
|
||||
"grad_norm": grad_norm,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank == 0:
|
||||
if rank <= 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss": loss,
|
||||
@@ -520,7 +520,13 @@ if __name__ == "__main__":
|
||||
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
|
||||
@@ -11,6 +11,7 @@ from torch.distributed.checkpoint.optimizer import load_sharded_optimizer_state_
|
||||
from torch.distributed.fsdp import FullOptimStateDictConfig, FullStateDictConfig
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.fsdp import StateDictType
|
||||
import dataclasses
|
||||
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
@@ -32,7 +33,7 @@ def save_checkpoint_optimizer(model, optimizer, rank, output_dir, step, discrimi
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank == 0 and not discriminator:
|
||||
if rank <= 0 and not discriminator:
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
config_dict = dict(model.config)
|
||||
@@ -44,13 +45,50 @@ def save_checkpoint_optimizer(model, optimizer, rank, output_dir, step, discrimi
|
||||
optimizer_path = os.path.join(save_dir, "optimizer.pt")
|
||||
torch.save(optim_state, optimizer_path)
|
||||
else:
|
||||
weight_path = os.path.join(save_dir, "discriminator_pytorch_model.safetensors")
|
||||
weight_path = os.path.join(save_dstate_dictir, "discriminator_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
optimizer_path = os.path.join(save_dir, "discriminator_optimizer.pt")
|
||||
torch.save(optim_state, optimizer_path)
|
||||
main_print(f"--> checkpoint saved at step {step}")
|
||||
|
||||
|
||||
def save_checkpoint_v1(transformer, rank, output_dir, step):
|
||||
# from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
# from torch.distributed.fsdp import StateDictType, FullStateDictConfig
|
||||
|
||||
# Configure FSDP to save full state dict
|
||||
FSDP.set_state_dict_type(
|
||||
transformer,
|
||||
state_dict_type=StateDictType.FULL_STATE_DICT,
|
||||
state_dict_config=FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
)
|
||||
|
||||
# Now get the state dict
|
||||
cpu_state = transformer.state_dict()
|
||||
|
||||
# Save it (only on rank 0 since we used rank0_only=True)
|
||||
# if torch.distributed.get_rank() == 0:
|
||||
# torch.save(state_dict, "model_checkpoint.pt")
|
||||
if rank <= 0:
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
# weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.pt")
|
||||
print(weight_path)
|
||||
# save_file(cpu_state, weight_path)
|
||||
torch.save(cpu_state, weight_path)
|
||||
config_dict = transformer.hf_config
|
||||
if "dtype" in config_dict:
|
||||
del config_dict["dtype"] # TODO
|
||||
config_path = os.path.join(save_dir, "config.json")
|
||||
# save dict as json
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
main_print(f"--> checkpoint saved at step {step}")
|
||||
|
||||
|
||||
|
||||
def save_checkpoint(transformer, rank, output_dir, step):
|
||||
main_print(f"--> saving checkpoint at step {step}")
|
||||
with FSDP.state_dict_type(
|
||||
@@ -60,7 +98,7 @@ def save_checkpoint(transformer, rank, output_dir, step):
|
||||
):
|
||||
cpu_state = transformer.state_dict()
|
||||
# todo move to get_state_dict
|
||||
if rank == 0:
|
||||
if rank <= 0:
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
@@ -98,7 +136,7 @@ def save_checkpoint_generator_discriminator(
|
||||
hf_weight_dir = os.path.join(save_dir, "hf_weights")
|
||||
os.makedirs(hf_weight_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank == 0:
|
||||
if rank <= 0:
|
||||
config_dict = dict(model.config)
|
||||
config_path = os.path.join(hf_weight_dir, "config.json")
|
||||
# save dict as json
|
||||
@@ -139,7 +177,7 @@ def save_checkpoint_generator_discriminator(
|
||||
optim_state = FSDP.optim_state_dict(discriminator, discriminator_optimizer)
|
||||
model_state = discriminator.state_dict()
|
||||
state_dict = {"optimizer": optim_state, "model": model_state}
|
||||
if rank == 0:
|
||||
if rank <= 0:
|
||||
discriminator_fsdp_state_fil = os.path.join(discriminator_fsdp_state_dir, "discriminator_state.pt")
|
||||
torch.save(state_dict, discriminator_fsdp_state_fil)
|
||||
|
||||
@@ -178,7 +216,7 @@ def load_full_state_model(model, optimizer, checkpoint_file, rank):
|
||||
):
|
||||
discriminator_state = torch.load(checkpoint_file)
|
||||
model_state = discriminator_state["model"]
|
||||
if rank == 0:
|
||||
if rank <= 0:
|
||||
optim_state = discriminator_state["optimizer"]
|
||||
else:
|
||||
optim_state = None
|
||||
@@ -241,7 +279,7 @@ def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step, pipelin
|
||||
optimizer,
|
||||
)
|
||||
|
||||
if rank == 0:
|
||||
if rank <= 0:
|
||||
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
import platform
|
||||
|
||||
import accelerate
|
||||
import peft
|
||||
import torch
|
||||
import transformers
|
||||
from transformers.utils import is_torch_cuda_available, is_torch_npu_available
|
||||
|
||||
VERSION = "1.2.0"
|
||||
|
||||
if __name__ == "__main__":
|
||||
info = {
|
||||
"FastVideo version": VERSION,
|
||||
"Platform": platform.platform(),
|
||||
"Python version": platform.python_version(),
|
||||
"PyTorch version": torch.__version__,
|
||||
"Transformers version": transformers.__version__,
|
||||
"Accelerate version": accelerate.__version__,
|
||||
"PEFT version": peft.__version__,
|
||||
}
|
||||
|
||||
if is_torch_cuda_available():
|
||||
info["PyTorch version"] += " (GPU)"
|
||||
info["GPU type"] = torch.cuda.get_device_name()
|
||||
|
||||
if is_torch_npu_available():
|
||||
info["PyTorch version"] += " (NPU)"
|
||||
info["NPU type"] = torch.npu.get_device_name()
|
||||
info["CANN version"] = torch.version.cann # codespell:ignore
|
||||
|
||||
try:
|
||||
import bitsandbytes
|
||||
|
||||
info["Bitsandbytes version"] = bitsandbytes.__version__
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
print("\n" + "\n".join([f"- {key}: {value}" for key, value in info.items()]) + "\n")
|
||||
@@ -1,88 +0,0 @@
|
||||
import torch
|
||||
|
||||
mochi_latents_mean = torch.tensor([
|
||||
-0.06730895953510081,
|
||||
-0.038011381506090416,
|
||||
-0.07477820912866141,
|
||||
-0.05565264470995561,
|
||||
0.012767231469026969,
|
||||
-0.04703542746246419,
|
||||
0.043896967884726704,
|
||||
-0.09346305707025976,
|
||||
-0.09918314763016893,
|
||||
-0.008729793427399178,
|
||||
-0.011931556316503654,
|
||||
-0.0321993391887285,
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
mochi_latents_std = torch.tensor([
|
||||
0.9263795028493863,
|
||||
0.9248894543193766,
|
||||
0.9393059390890617,
|
||||
0.959253732819592,
|
||||
0.8244560132752793,
|
||||
0.917259975397747,
|
||||
0.9294154431013696,
|
||||
1.3720942357788521,
|
||||
0.881393668867029,
|
||||
0.9168315692124348,
|
||||
0.9185249279345552,
|
||||
0.9274757570805041,
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
mochi_scaling_factor = 1.0
|
||||
|
||||
|
||||
wan_latents_mean = torch.tensor([
|
||||
-0.7571,
|
||||
-0.7089,
|
||||
-0.9113,
|
||||
0.1075,
|
||||
-0.1745,
|
||||
0.9653,
|
||||
-0.1517,
|
||||
1.5508,
|
||||
0.4134,
|
||||
-0.0715,
|
||||
0.5517,
|
||||
-0.3632,
|
||||
-0.1922,
|
||||
-0.9497,
|
||||
0.2503,
|
||||
-0.2921,
|
||||
]).view(1, 16, 1, 1, 1)
|
||||
wan_latents_std = torch.tensor([
|
||||
2.8184,
|
||||
1.4541,
|
||||
2.3275,
|
||||
2.6558,
|
||||
1.2196,
|
||||
1.7708,
|
||||
2.6052,
|
||||
2.0743,
|
||||
3.2687,
|
||||
2.1526,
|
||||
2.8652,
|
||||
1.5579,
|
||||
1.6382,
|
||||
1.1253,
|
||||
2.8251,
|
||||
1.916,
|
||||
]).view(1, 16, 1, 1, 1)
|
||||
|
||||
|
||||
def normalize_dit_input(model_type, latents):
|
||||
if model_type == "mochi":
|
||||
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
|
||||
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
|
||||
latents = (latents - latents_mean) / latents_std
|
||||
return latents
|
||||
elif model_type == "hunyuan_hf":
|
||||
return latents * 0.476986
|
||||
elif model_type == "hunyuan":
|
||||
return latents * 0.476986
|
||||
elif model_type == "wan":
|
||||
latents_mean = wan_latents_mean.to(latents.device, latents.dtype)
|
||||
latents_std = wan_latents_std.to(latents.device, latents.dtype)
|
||||
latents = (latents - latents_mean) / latents_std
|
||||
return latents
|
||||
else:
|
||||
raise NotImplementedError(f"model_type {model_type} not supported")
|
||||
+2
-69
@@ -3,9 +3,9 @@ from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi, AutoencoderKLWan
|
||||
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi
|
||||
from torch import nn
|
||||
from transformers import AutoTokenizer, T5EncoderModel, UMT5EncoderModel
|
||||
from transformers import AutoTokenizer, T5EncoderModel
|
||||
|
||||
from fastvideo.models.hunyuan.modules.models import (HYVideoDiffusionTransformer, MMDoubleStreamBlock,
|
||||
MMSingleStreamBlock)
|
||||
@@ -14,7 +14,6 @@ from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import AutoencoderKLC
|
||||
from fastvideo.models.hunyuan_hf.modeling_hunyuan import (HunyuanVideoSingleTransformerBlock,
|
||||
HunyuanVideoTransformer3DModel, HunyuanVideoTransformerBlock)
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel, MochiTransformerBlock
|
||||
from fastvideo.models.wan_hf.modeling_wan import WanTransformer3DModel, WanTransformerBlock
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
hunyuan_config = {
|
||||
@@ -201,48 +200,6 @@ class MochiTextEncoderWrapper(nn.Module):
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
class WanTextEncoderWrapper(nn.Module):
|
||||
|
||||
def __init__(self, pretrained_model_name_or_path, device):
|
||||
super().__init__()
|
||||
self.text_encoder = UMT5EncoderModel.from_pretrained(os.path.join(pretrained_model_name_or_path,
|
||||
"text_encoder")).to(device)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(os.path.join(pretrained_model_name_or_path, "tokenizer"))
|
||||
self.max_sequence_length = 256
|
||||
|
||||
def encode_prompt(self, prompt):
|
||||
device = self.text_encoder.device
|
||||
dtype = self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=self.max_sequence_length,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
prompt_attention_mask = text_inputs.attention_mask
|
||||
prompt_attention_mask = prompt_attention_mask.bool().to(device)
|
||||
|
||||
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
|
||||
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.max_sequence_length - 1:-1])
|
||||
main_print(f"Truncated text input: {prompt} to: {removed_text} for model input.")
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.view(batch_size, seq_len, -1)
|
||||
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
def load_hunyuan_state_dict(model, dit_model_name_or_path):
|
||||
load_key = "module"
|
||||
@@ -283,20 +240,6 @@ def load_transformer(
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
elif model_type == "wan":
|
||||
if dit_model_name_or_path:
|
||||
transformer = WanTransformer3DModel.from_pretrained(
|
||||
dit_model_name_or_path,
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
else:
|
||||
transformer = WanTransformer3DModel.from_pretrained(
|
||||
pretrained_model_name_or_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
elif model_type == "hunyuan_hf":
|
||||
if dit_model_name_or_path:
|
||||
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
|
||||
@@ -340,12 +283,6 @@ def load_vae(model_type, pretrained_model_name_or_path):
|
||||
torch_dtype=weight_dtype).to("cuda")
|
||||
autocast_type = torch.bfloat16
|
||||
fps = 24
|
||||
elif model_type == "wan":
|
||||
vae = AutoencoderKLWan.from_pretrained(pretrained_model_name_or_path,
|
||||
subfolder="vae",
|
||||
torch_dtype=weight_dtype).to("cuda")
|
||||
autocast_type = torch.bfloat16
|
||||
fps = 24
|
||||
elif model_type == "hunyuan":
|
||||
vae_precision = torch.float32
|
||||
vae_path = os.path.join(pretrained_model_name_or_path, "hunyuan-video-t2v-720p/vae")
|
||||
@@ -374,8 +311,6 @@ def load_vae(model_type, pretrained_model_name_or_path):
|
||||
def load_text_encoder(model_type, pretrained_model_name_or_path, device):
|
||||
if model_type == "mochi":
|
||||
text_encoder = MochiTextEncoderWrapper(pretrained_model_name_or_path, device)
|
||||
elif model_type == "wan":
|
||||
text_encoder = WanTextEncoderWrapper(pretrained_model_name_or_path, device)
|
||||
elif model_type == "hunyuan" or "hunyuan_hf":
|
||||
text_encoder = HunyuanTextEncoderWrapper(pretrained_model_name_or_path, device)
|
||||
else:
|
||||
@@ -387,8 +322,6 @@ def get_no_split_modules(transformer):
|
||||
# if of type MochiTransformer3DModel
|
||||
if isinstance(transformer, MochiTransformer3DModel):
|
||||
return (MochiTransformerBlock, )
|
||||
elif isinstance(transformer, WanTransformer3DModel):
|
||||
return (WanTransformerBlock, )
|
||||
elif isinstance(transformer, HunyuanVideoTransformer3DModel):
|
||||
return (HunyuanVideoSingleTransformerBlock, HunyuanVideoTransformerBlock)
|
||||
elif isinstance(transformer, HYVideoDiffusionTransformer):
|
||||
|
||||
@@ -129,22 +129,13 @@ def sample_validation_video(
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latent_model_input.shape[0])
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
if model_type == "wan":
|
||||
pred_kwargs = {
|
||||
"hidden_states": latent_model_input,
|
||||
"encoder_hidden_states": prompt_embeds,
|
||||
"timestep":timestep,
|
||||
"return_dict":False,
|
||||
}
|
||||
else:
|
||||
pred_kwargs = {
|
||||
"hidden_states": latent_model_input,
|
||||
"encoder_hidden_states": prompt_embeds,
|
||||
"timestep":timestep,
|
||||
"encoder_attention_mask":prompt_attention_mask,
|
||||
"return_dict":False,
|
||||
}
|
||||
noise_pred = transformer(**pred_kwargs)[0]
|
||||
noise_pred = transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# Mochi CFG + Sampling runs in FP32
|
||||
noise_pred = noise_pred.to(torch.float32)
|
||||
@@ -175,12 +166,10 @@ def sample_validation_video(
|
||||
# denormalize with the mean and std if available and not None
|
||||
has_latents_mean = (hasattr(vae.config, "latents_mean") and vae.config.latents_mean is not None)
|
||||
has_latents_std = (hasattr(vae.config, "latents_std") and vae.config.latents_std is not None)
|
||||
if model_type == "wan":
|
||||
vae.config.scaling_factor = 1
|
||||
if has_latents_mean and has_latents_std:
|
||||
latents_mean = (torch.tensor(vae.config.latents_mean).view(1, num_channels_latents, 1, 1,
|
||||
latents_mean = (torch.tensor(vae.config.latents_mean).view(1, 12, 1, 1,
|
||||
1).to(latents.device, latents.dtype))
|
||||
latents_std = (torch.tensor(vae.config.latents_std).view(1, num_channels_latents, 1, 1, 1).to(latents.device, latents.dtype))
|
||||
latents_std = (torch.tensor(vae.config.latents_std).view(1, 12, 1, 1, 1).to(latents.device, latents.dtype))
|
||||
latents = latents * latents_std / vae.config.scaling_factor + latents_mean
|
||||
else:
|
||||
latents = latents / vae.config.scaling_factor
|
||||
@@ -213,15 +202,14 @@ def log_validation(
|
||||
vae_spatial_scale_factor = 8
|
||||
vae_temporal_scale_factor = 6
|
||||
num_channels_latents = 12
|
||||
elif args.model_type == "hunyuan" or "hunyuan_hf" or "wan":
|
||||
elif args.model_type == "hunyuan" or "hunyuan_hf":
|
||||
vae_spatial_scale_factor = 8
|
||||
vae_temporal_scale_factor = 4
|
||||
num_channels_latents = 16
|
||||
else:
|
||||
raise ValueError(f"Model type {args.model_type} not supported")
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.pretrained_model_name_or_path)
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
vae.enable_tiling()
|
||||
if scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler(shift=shift)
|
||||
else:
|
||||
|
||||
@@ -1,420 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
from collections import defaultdict
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def configure_sta(mode: str = 'STA_searching',
|
||||
layer_num: int = 40,
|
||||
time_step_num: int = 50,
|
||||
head_num: int = 40,
|
||||
**kwargs) -> List[List[List[Any]]]:
|
||||
"""
|
||||
Configure Sliding Tile Attention (STA) parameters based on the specified mode.
|
||||
|
||||
Parameters:
|
||||
----------
|
||||
mode : str
|
||||
The STA mode to use. Options are:
|
||||
- 'STA_searching': Generate a set of mask candidates for initial search
|
||||
- 'STA_tuning': Select best mask strategy based on previously saved results
|
||||
- 'STA_inference': Load and use a previously tuned mask strategy
|
||||
layer_num: int, number of layers
|
||||
time_step_num: int, number of timesteps
|
||||
head_num: int, number of heads
|
||||
|
||||
**kwargs : dict
|
||||
Mode-specific parameters:
|
||||
|
||||
For 'STA_searching':
|
||||
- mask_candidates: list of str, optional, mask candidates to use
|
||||
- mask_selected: list of int, optional, indices of selected masks
|
||||
|
||||
For 'STA_tuning':
|
||||
- mask_search_files_path: str, required, path to mask search results
|
||||
- mask_candidates: list of str, optional, mask candidates to use
|
||||
- mask_selected: list of int, optional, indices of selected masks
|
||||
- skip_time_steps: int, optional, number of time steps to use full attention (default 12)
|
||||
- save_dir: str, optional, directory to save mask strategy (default "mask_candidates")
|
||||
|
||||
For 'STA_inference':
|
||||
- load_path: str, optional, path to load mask strategy (default "mask_candidates/mask_strategy.json")
|
||||
"""
|
||||
valid_modes = [
|
||||
'STA_searching', 'STA_tuning', 'STA_inference', 'STA_tuning_cfg'
|
||||
]
|
||||
if mode not in valid_modes:
|
||||
raise ValueError(f"Mode must be one of {valid_modes}, got {mode}")
|
||||
|
||||
if mode == 'STA_searching':
|
||||
# Get parameters with defaults
|
||||
mask_candidates: Optional[List[str]] = kwargs.get('mask_candidates')
|
||||
if mask_candidates is None:
|
||||
raise ValueError(
|
||||
"mask_candidates is required for STA_searching mode")
|
||||
mask_selected: List[int] = kwargs.get('mask_selected',
|
||||
list(range(len(mask_candidates))))
|
||||
|
||||
# Parse selected masks
|
||||
selected_masks: List[List[int]] = []
|
||||
for index in mask_selected:
|
||||
mask = mask_candidates[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks.append(masks_list)
|
||||
|
||||
# Create 3D mask structure with fixed dimensions (t=50, l=60)
|
||||
masks_3d: List[List[List[List[int]]]] = []
|
||||
for i in range(time_step_num): # Fixed t dimension = 50
|
||||
row = []
|
||||
for j in range(layer_num): # Fixed l dimension = 60
|
||||
row.append(selected_masks) # Add all masks at each position
|
||||
masks_3d.append(row)
|
||||
|
||||
return masks_3d
|
||||
|
||||
elif mode == 'STA_tuning':
|
||||
# Get required parameters
|
||||
mask_search_files_path: Optional[str] = kwargs.get(
|
||||
'mask_search_files_path')
|
||||
if not mask_search_files_path:
|
||||
raise ValueError(
|
||||
"mask_search_files_path is required for STA_tuning mode")
|
||||
|
||||
# Get optional parameters with defaults
|
||||
mask_candidates_tuning: Optional[List[str]] = kwargs.get(
|
||||
'mask_candidates')
|
||||
if mask_candidates_tuning is None:
|
||||
raise ValueError("mask_candidates is required for STA_tuning mode")
|
||||
mask_selected_tuning: List[int] = kwargs.get(
|
||||
'mask_selected', list(range(len(mask_candidates_tuning))))
|
||||
skip_time_steps_tuning: Optional[int] = kwargs.get('skip_time_steps')
|
||||
save_dir_tuning: Optional[str] = kwargs.get('save_dir',
|
||||
"mask_candidates")
|
||||
|
||||
# Parse selected masks
|
||||
selected_masks_tuning: List[List[int]] = []
|
||||
for index in mask_selected_tuning:
|
||||
mask = mask_candidates_tuning[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks_tuning.append(masks_list)
|
||||
|
||||
# Read JSON results
|
||||
results = read_specific_json_files(mask_search_files_path)
|
||||
averaged_results = average_head_losses(results, selected_masks_tuning)
|
||||
|
||||
# Add full attention mask for specific cases
|
||||
full_attention_mask_tuning: Optional[List[int]] = kwargs.get(
|
||||
'full_attention_mask')
|
||||
if full_attention_mask_tuning is not None:
|
||||
selected_masks_tuning.append(full_attention_mask_tuning)
|
||||
|
||||
# Select best mask strategy
|
||||
timesteps_tuning: int = kwargs.get('timesteps', time_step_num)
|
||||
if skip_time_steps_tuning is None:
|
||||
skip_time_steps_tuning = 12
|
||||
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
|
||||
averaged_results, selected_masks_tuning, skip_time_steps_tuning,
|
||||
timesteps_tuning, head_num)
|
||||
|
||||
# Save mask strategy
|
||||
if save_dir_tuning is not None:
|
||||
os.makedirs(save_dir_tuning, exist_ok=True)
|
||||
file_path = os.path.join(
|
||||
save_dir_tuning,
|
||||
f'mask_strategy_s{skip_time_steps_tuning}.json')
|
||||
with open(file_path, 'w') as f:
|
||||
json.dump(mask_strategy, f, indent=4)
|
||||
print(f"Successfully saved mask_strategy to {file_path}")
|
||||
|
||||
# Print sparsity and strategy counts for information
|
||||
print(f"Overall sparsity: {sparsity:.4f}")
|
||||
print("\nStrategy usage counts:")
|
||||
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
|
||||
for strategy, count in strategy_counts.items():
|
||||
print(
|
||||
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
|
||||
)
|
||||
|
||||
# Convert dictionary to 3D list with fixed dimensions
|
||||
mask_strategy_3d = dict_to_3d_list(mask_strategy,
|
||||
t_max=time_step_num,
|
||||
l_max=layer_num,
|
||||
h_max=head_num)
|
||||
|
||||
return mask_strategy_3d
|
||||
elif mode == 'STA_tuning_cfg':
|
||||
# Get required parameters for both positive and negative paths
|
||||
mask_search_files_path_pos: Optional[str] = kwargs.get(
|
||||
'mask_search_files_path_pos')
|
||||
mask_search_files_path_neg: Optional[str] = kwargs.get(
|
||||
'mask_search_files_path_neg')
|
||||
save_dir_cfg: Optional[str] = kwargs.get('save_dir')
|
||||
|
||||
if not mask_search_files_path_pos or not mask_search_files_path_neg or not save_dir_cfg:
|
||||
raise ValueError(
|
||||
"mask_search_files_path_pos, mask_search_files_path_neg, and save_dir are required for STA_tuning_cfg mode"
|
||||
)
|
||||
|
||||
# Get optional parameters with defaults
|
||||
mask_candidates_cfg: Optional[List[str]] = kwargs.get('mask_candidates')
|
||||
if mask_candidates_cfg is None:
|
||||
raise ValueError(
|
||||
"mask_candidates is required for STA_tuning_cfg mode")
|
||||
mask_selected_cfg: List[int] = kwargs.get(
|
||||
'mask_selected', list(range(len(mask_candidates_cfg))))
|
||||
skip_time_steps_cfg: Optional[int] = kwargs.get('skip_time_steps')
|
||||
|
||||
# Parse selected masks
|
||||
selected_masks_cfg: List[List[int]] = []
|
||||
for index in mask_selected_cfg:
|
||||
mask = mask_candidates_cfg[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks_cfg.append(masks_list)
|
||||
|
||||
# Read JSON results for both positive and negative paths
|
||||
pos_results = read_specific_json_files(mask_search_files_path_pos)
|
||||
neg_results = read_specific_json_files(mask_search_files_path_neg)
|
||||
# Combine positive and negative results into one list
|
||||
combined_results = pos_results + neg_results
|
||||
|
||||
# Average the combined results
|
||||
averaged_results = average_head_losses(combined_results,
|
||||
selected_masks_cfg)
|
||||
|
||||
# Add full attention mask for specific cases
|
||||
full_attention_mask_cfg: Optional[List[int]] = kwargs.get(
|
||||
'full_attention_mask')
|
||||
if full_attention_mask_cfg is not None:
|
||||
selected_masks_cfg.append(full_attention_mask_cfg)
|
||||
|
||||
timesteps_cfg: int = kwargs.get('timesteps', time_step_num)
|
||||
if skip_time_steps_cfg is None:
|
||||
skip_time_steps_cfg = 12
|
||||
# Select best mask strategy using combined results
|
||||
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
|
||||
averaged_results, selected_masks_cfg, skip_time_steps_cfg,
|
||||
timesteps_cfg, head_num)
|
||||
|
||||
# Save mask strategy
|
||||
os.makedirs(save_dir_cfg, exist_ok=True)
|
||||
file_path = os.path.join(save_dir_cfg,
|
||||
f'mask_strategy_s{skip_time_steps_cfg}.json')
|
||||
with open(file_path, 'w') as f:
|
||||
json.dump(mask_strategy, f, indent=4)
|
||||
print(f"Successfully saved mask_strategy to {file_path}")
|
||||
|
||||
# Print sparsity and strategy counts for information
|
||||
print(f"Overall sparsity: {sparsity:.4f}")
|
||||
print("\nStrategy usage counts:")
|
||||
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
|
||||
for strategy, count in strategy_counts.items():
|
||||
print(
|
||||
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
|
||||
)
|
||||
|
||||
# Convert dictionary to 3D list with fixed dimensions
|
||||
mask_strategy_3d = dict_to_3d_list(mask_strategy,
|
||||
t_max=time_step_num,
|
||||
l_max=layer_num,
|
||||
h_max=head_num)
|
||||
|
||||
return mask_strategy_3d
|
||||
|
||||
else: # STA_inference
|
||||
# Get parameters with defaults
|
||||
load_path: Optional[str] = kwargs.get(
|
||||
'load_path', "mask_candidates/mask_strategy.json")
|
||||
if load_path is None:
|
||||
raise ValueError("load_path is required for STA_inference mode")
|
||||
|
||||
# Load previously saved mask strategy
|
||||
with open(load_path) as f:
|
||||
mask_strategy = json.load(f)
|
||||
|
||||
# Convert dictionary to 3D list with fixed dimensions
|
||||
mask_strategy_3d = dict_to_3d_list(mask_strategy,
|
||||
t_max=time_step_num,
|
||||
l_max=layer_num,
|
||||
h_max=head_num)
|
||||
|
||||
return mask_strategy_3d
|
||||
|
||||
|
||||
# Helper functions
|
||||
|
||||
|
||||
def read_specific_json_files(folder_path: str) -> List[Dict[str, Any]]:
|
||||
"""Read and parse JSON files containing mask search results."""
|
||||
json_contents: List[Dict[str, Any]] = []
|
||||
|
||||
# List files only in the current directory (no walk)
|
||||
files = os.listdir(folder_path)
|
||||
# Filter files
|
||||
matching_files = [f for f in files if 'mask' in f and f.endswith('.json')]
|
||||
print(f"Found {len(matching_files)} matching files: {matching_files}")
|
||||
|
||||
for file_name in matching_files:
|
||||
file_path = os.path.join(folder_path, file_name)
|
||||
with open(file_path) as file:
|
||||
data = json.load(file)
|
||||
json_contents.append(data)
|
||||
|
||||
return json_contents
|
||||
|
||||
|
||||
def average_head_losses(
|
||||
results: List[Dict[str, Any]],
|
||||
selected_masks: List[List[int]]) -> Dict[str, Dict[str, np.ndarray]]:
|
||||
"""Average losses across all prompts for each mask strategy."""
|
||||
# Initialize a dictionary to store the averaged results
|
||||
averaged_losses: Dict[str, Dict[str, np.ndarray]] = {}
|
||||
loss_type = 'L2_loss'
|
||||
# Get all loss types (e.g., 'L2_loss')
|
||||
averaged_losses[loss_type] = {}
|
||||
|
||||
for mask in selected_masks:
|
||||
mask_str = str(mask)
|
||||
data_shape = np.array(results[0][loss_type][mask_str]).shape
|
||||
accumulated_data = np.zeros(data_shape)
|
||||
|
||||
# Sum across all prompts
|
||||
for prompt_result in results:
|
||||
accumulated_data += np.array(prompt_result[loss_type][mask_str])
|
||||
|
||||
# Average by dividing by number of prompts
|
||||
averaged_data = accumulated_data / len(results)
|
||||
averaged_losses[loss_type][mask_str] = averaged_data
|
||||
|
||||
return averaged_losses
|
||||
|
||||
|
||||
def select_best_mask_strategy(
|
||||
averaged_results: Dict[str, Dict[str, np.ndarray]],
|
||||
selected_masks: List[List[int]],
|
||||
skip_time_steps: int = 12,
|
||||
timesteps: int = 50,
|
||||
head_num: int = 40
|
||||
) -> Tuple[Dict[str, List[int]], float, Dict[str, int]]:
|
||||
"""Select the best mask strategy for each head based on loss minimization."""
|
||||
best_mask_strategy: Dict[str, List[int]] = {}
|
||||
loss_type = 'L2_loss'
|
||||
# Get the shape of time steps and layers
|
||||
layers = len(averaged_results[loss_type][str(selected_masks[0])][0])
|
||||
|
||||
# Counter for sparsity calculation
|
||||
total_tokens = 0 # total number of masked tokens
|
||||
total_length = 0 # total sequence length
|
||||
|
||||
strategy_counts: Dict[str, int] = {
|
||||
str(strategy): 0
|
||||
for strategy in selected_masks
|
||||
}
|
||||
full_attn_strategy = selected_masks[-1] # Last strategy is full attention
|
||||
print(f"Strategy {full_attn_strategy}, skip first {skip_time_steps} steps ")
|
||||
|
||||
for t in range(timesteps):
|
||||
for layer_idx in range(layers):
|
||||
for h in range(head_num):
|
||||
if t < skip_time_steps: # First steps use full attention
|
||||
strategy = full_attn_strategy
|
||||
else:
|
||||
# Get losses for this head across all strategies
|
||||
head_losses = []
|
||||
for strategy in selected_masks[:
|
||||
-1]: # Exclude full attention
|
||||
head_losses.append(averaged_results[loss_type][str(
|
||||
strategy)][t][layer_idx][h])
|
||||
|
||||
# Find which strategy gives minimum loss
|
||||
best_strategy_idx = np.argmin(head_losses)
|
||||
strategy = selected_masks[best_strategy_idx]
|
||||
|
||||
best_mask_strategy[f'{t}_{layer_idx}_{h}'] = strategy
|
||||
|
||||
# Calculate sparsity
|
||||
nums = strategy # strategy is already a list of numbers
|
||||
total_tokens += nums[0] * nums[1] * nums[
|
||||
2] # masked tokens for chosen strategy
|
||||
total_length += full_attn_strategy[0] * full_attn_strategy[
|
||||
1] * full_attn_strategy[2]
|
||||
|
||||
# Count strategy usage
|
||||
strategy_counts[str(strategy)] += 1
|
||||
|
||||
overall_sparsity = 1 - total_tokens / total_length
|
||||
|
||||
return best_mask_strategy, overall_sparsity, strategy_counts
|
||||
|
||||
|
||||
def dict_to_3d_list(mask_strategy: Optional[Dict[str, List[int]]],
|
||||
t_max: int = 50,
|
||||
l_max: int = 60,
|
||||
h_max: int = 24) -> List[List[List[Optional[List[int]]]]]:
|
||||
result: List[List[List[Optional[List[int]]]]] = [[[
|
||||
None for _ in range(h_max)
|
||||
] for _ in range(l_max)] for _ in range(t_max)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
for key, value in mask_strategy.items():
|
||||
t, layer_idx, h = map(int, key.split('_'))
|
||||
result[t][layer_idx][h] = value
|
||||
return result
|
||||
|
||||
|
||||
def save_mask_search_results(
|
||||
mask_search_final_result: List[Dict[str, List[float]]],
|
||||
prompt: str,
|
||||
mask_strategies: List[str],
|
||||
output_dir: str = 'output/mask_search_result/') -> Optional[str]:
|
||||
if not mask_search_final_result:
|
||||
print("No mask search results to save")
|
||||
return None
|
||||
|
||||
# Create result dictionary with defaultdict for nested lists
|
||||
mask_search_dict: Dict[str, Dict[str, List[List[float]]]] = {
|
||||
"L2_loss": defaultdict(list),
|
||||
"L1_loss": defaultdict(list)
|
||||
}
|
||||
|
||||
mask_selected = list(range(len(mask_strategies)))
|
||||
selected_masks: List[List[int]] = []
|
||||
for index in mask_selected:
|
||||
mask = mask_strategies[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks.append(masks_list)
|
||||
|
||||
# Process each mask strategy
|
||||
for i, mask_strategy in enumerate(selected_masks):
|
||||
mask_strategy_str = str(mask_strategy)
|
||||
# Process L2 loss
|
||||
step_results: List[List[float]] = []
|
||||
for step_data in mask_search_final_result:
|
||||
if isinstance(step_data, dict) and "L2_loss" in step_data:
|
||||
layer_losses = [float(loss) for loss in step_data["L2_loss"]]
|
||||
step_results.append(layer_losses)
|
||||
mask_search_dict["L2_loss"][mask_strategy_str] = step_results
|
||||
|
||||
step_results = []
|
||||
for step_data in mask_search_final_result:
|
||||
if isinstance(step_data, dict) and "L1_loss" in step_data:
|
||||
layer_losses = [float(loss) for loss in step_data["L1_loss"]]
|
||||
step_results.append(layer_losses)
|
||||
mask_search_dict["L1_loss"][mask_strategy_str] = step_results
|
||||
|
||||
# Create the output directory if it doesn't exist
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
# Create a filename based on the first 20 characters of the prompt
|
||||
filename = prompt[:50].replace(" ", "_")
|
||||
filepath = os.path.join(output_dir, f'mask_search_{filename}.json')
|
||||
|
||||
# Save the results to a JSON file
|
||||
with open(filepath, 'w') as f:
|
||||
json.dump(mask_search_dict, f, indent=4)
|
||||
|
||||
print(f"Successfully saved mask research results to {filepath}")
|
||||
|
||||
return filepath
|
||||
@@ -3,14 +3,11 @@
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.attention.layer import (DistributedAttention,
|
||||
DistributedAttention_VSA,
|
||||
LocalAttention)
|
||||
from fastvideo.v1.attention.layer import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.attention.selector import get_attn_backend
|
||||
|
||||
__all__ = [
|
||||
"DistributedAttention",
|
||||
"DistributedAttention_VSA",
|
||||
"LocalAttention",
|
||||
"AttentionBackend",
|
||||
"AttentionMetadata",
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Type
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
@@ -14,7 +13,6 @@ from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
@@ -22,9 +20,7 @@ logger = init_logger(__name__)
|
||||
|
||||
|
||||
# TODO(will-refactor): move this to a utils file
|
||||
def dict_to_3d_list(
|
||||
mask_strategy: Dict[str,
|
||||
Any]) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
def dict_to_3d_list(mask_strategy) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
|
||||
|
||||
max_timesteps_idx = max(
|
||||
@@ -46,14 +42,14 @@ def dict_to_3d_list(
|
||||
|
||||
class RangeDict(dict):
|
||||
|
||||
def __getitem__(self, item: int) -> str:
|
||||
def __getitem__(self, item):
|
||||
for key in self.keys():
|
||||
if isinstance(key, tuple):
|
||||
low, high = key
|
||||
if low <= item <= high:
|
||||
return str(super().__getitem__(key))
|
||||
return super().__getitem__(key)
|
||||
elif key == item:
|
||||
return str(super().__getitem__(key))
|
||||
return super().__getitem__(key)
|
||||
raise KeyError(f"seq_len {item} not supported for STA")
|
||||
|
||||
|
||||
@@ -86,8 +82,6 @@ class SlidingTileAttentionBackend(AttentionBackend):
|
||||
@dataclass
|
||||
class SlidingTileAttentionMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
STA_param: List[List[
|
||||
Any]] # each timestep with one metadata, shape [num_layers, num_heads]
|
||||
|
||||
|
||||
class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
@@ -104,12 +98,8 @@ class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
forward_batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> SlidingTileAttentionMetadata:
|
||||
param = forward_batch.STA_param
|
||||
if param is None:
|
||||
return SlidingTileAttentionMetadata(
|
||||
current_timestep=current_timestep, STA_param=[])
|
||||
return SlidingTileAttentionMetadata(current_timestep=current_timestep,
|
||||
STA_param=param[current_timestep])
|
||||
|
||||
return SlidingTileAttentionMetadata(current_timestep=current_timestep, )
|
||||
|
||||
|
||||
class SlidingTileAttentionImpl(AttentionImpl):
|
||||
@@ -130,17 +120,17 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
if config_file is None:
|
||||
raise ValueError("FASTVIDEO_ATTENTION_CONFIG is not set")
|
||||
|
||||
# TODO(kevin): get mask strategy for different STA modes
|
||||
with open(config_file) as f:
|
||||
mask_strategy = json.load(f)
|
||||
self.mask_strategy = dict_to_3d_list(mask_strategy)
|
||||
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
|
||||
# STA config
|
||||
self.STA_base_tile_size = [6, 8, 8]
|
||||
self.dit_seq_shape_mapping = RangeDict({
|
||||
self.img_latent_shape_mapping = RangeDict({
|
||||
(115200, 115456): '30x48x80',
|
||||
82944: '36x48x48',
|
||||
69120: '18x48x80',
|
||||
@@ -155,9 +145,9 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
x = rearrange(x,
|
||||
"b (sp t h w) head d -> b (t sp h w) head d",
|
||||
sp=self.sp_size,
|
||||
t=self.dit_seq_shape_int[0] // self.sp_size,
|
||||
h=self.dit_seq_shape_int[1],
|
||||
w=self.dit_seq_shape_int[2])
|
||||
t=self.img_latent_shape_int[0] // self.sp_size,
|
||||
h=self.img_latent_shape_int[1],
|
||||
w=self.img_latent_shape_int[2])
|
||||
return rearrange(
|
||||
x,
|
||||
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
|
||||
@@ -181,9 +171,9 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
return rearrange(x,
|
||||
"b (t sp h w) head d -> b (sp t h w) head d",
|
||||
sp=self.sp_size,
|
||||
t=self.dit_seq_shape_int[0] // self.sp_size,
|
||||
h=self.dit_seq_shape_int[1],
|
||||
w=self.dit_seq_shape_int[2])
|
||||
t=self.img_latent_shape_int[0] // self.sp_size,
|
||||
h=self.img_latent_shape_int[1],
|
||||
w=self.img_latent_shape_int[2])
|
||||
|
||||
def preprocess_qkv(
|
||||
self,
|
||||
@@ -191,12 +181,14 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
img_sequence_length = qkv.shape[1]
|
||||
self.dit_seq_shape_str = self.dit_seq_shape_mapping[img_sequence_length]
|
||||
self.full_window_size = self.full_window_mapping[self.dit_seq_shape_str]
|
||||
self.dit_seq_shape_int = list(
|
||||
map(int, self.dit_seq_shape_str.split('x')))
|
||||
self.img_seq_length = self.dit_seq_shape_int[
|
||||
0] * self.dit_seq_shape_int[1] * self.dit_seq_shape_int[2]
|
||||
self.img_latent_shape_str = self.img_latent_shape_mapping[
|
||||
img_sequence_length]
|
||||
self.full_window_size = self.full_window_mapping[
|
||||
self.img_latent_shape_str]
|
||||
self.img_latent_shape_int = list(
|
||||
map(int, self.img_latent_shape_str.split('x')))
|
||||
self.img_seq_length = self.img_latent_shape_int[
|
||||
0] * self.img_latent_shape_int[1] * self.img_latent_shape_int[2]
|
||||
return self.tile(qkv)
|
||||
|
||||
def postprocess_output(
|
||||
@@ -213,24 +205,16 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
v: torch.Tensor,
|
||||
attn_metadata: SlidingTileAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
if self.mask_strategy is None:
|
||||
raise ValueError(
|
||||
"mask_strategy cannot be None for SlidingTileAttention")
|
||||
if self.mask_strategy[0] is None:
|
||||
raise ValueError(
|
||||
"mask_strategy[0] cannot be None for SlidingTileAttention")
|
||||
|
||||
assert self.mask_strategy is not None, "mask_strategy cannot be None for SlidingTileAttention"
|
||||
assert self.mask_strategy[
|
||||
0] is not None, "mask_strategy[0] cannot be None for SlidingTileAttention"
|
||||
|
||||
timestep = attn_metadata.current_timestep
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
forward_batch = forward_context.forward_batch
|
||||
if forward_batch is None:
|
||||
raise ValueError("forward_batch cannot be None")
|
||||
# pattern:'.double_blocks.0.attn.impl' or '.single_blocks.0.attn.impl'
|
||||
layer_idx = int(self.prefix.split('.')[-3])
|
||||
if attn_metadata.STA_param is None or len(
|
||||
attn_metadata.STA_param) <= layer_idx:
|
||||
raise ValueError("Invalid STA_param")
|
||||
STA_param = attn_metadata.STA_param[layer_idx]
|
||||
|
||||
# TODO: remove hardcode
|
||||
|
||||
text_length = q.shape[1] - self.img_seq_length
|
||||
has_text = text_length > 0
|
||||
@@ -243,56 +227,15 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
sp_group = get_sp_group()
|
||||
current_rank = sp_group.rank_in_group
|
||||
start_head = current_rank * head_num
|
||||
|
||||
# searching or tuning mode
|
||||
if len(STA_param) < head_num * sp_group.world_size:
|
||||
sparse_attn_hidden_states_all = []
|
||||
full_mask_window = STA_param[-1]
|
||||
for window_size in STA_param[:-1]:
|
||||
sparse_hidden_states = sliding_tile_attention(
|
||||
query, key, value, [window_size] * head_num, text_length,
|
||||
has_text, self.dit_seq_shape_str).transpose(1, 2)
|
||||
sparse_attn_hidden_states_all.append(sparse_hidden_states)
|
||||
|
||||
hidden_states = sliding_tile_attention(
|
||||
query, key, value, [full_mask_window] * head_num, text_length,
|
||||
has_text, self.dit_seq_shape_str).transpose(1, 2)
|
||||
|
||||
attn_L2_loss = []
|
||||
attn_L1_loss = []
|
||||
# average loss across all heads
|
||||
for sparse_attn_hidden_states in sparse_attn_hidden_states_all:
|
||||
# L2 loss
|
||||
attn_L2_loss_ = torch.mean((sparse_attn_hidden_states.float() -
|
||||
hidden_states.float())**2,
|
||||
dim=[0, 1, 3]).cpu().numpy()
|
||||
attn_L2_loss_ = [round(float(x), 6) for x in attn_L2_loss_]
|
||||
attn_L2_loss.append(attn_L2_loss_)
|
||||
# L1 loss
|
||||
attn_L1_loss_ = torch.mean(
|
||||
torch.abs(sparse_attn_hidden_states.float() -
|
||||
hidden_states.float()),
|
||||
dim=[0, 1, 3]).cpu().numpy()
|
||||
attn_L1_loss_ = [round(float(x), 6) for x in attn_L1_loss_]
|
||||
attn_L1_loss.append(attn_L1_loss_)
|
||||
|
||||
layer_loss_save = {"L2_loss": attn_L2_loss, "L1_loss": attn_L1_loss}
|
||||
|
||||
if forward_batch.is_cfg_negative:
|
||||
if forward_batch.mask_search_final_result_neg is not None:
|
||||
forward_batch.mask_search_final_result_neg[timestep].append(
|
||||
layer_loss_save)
|
||||
else:
|
||||
if forward_batch.mask_search_final_result_pos is not None:
|
||||
forward_batch.mask_search_final_result_pos[timestep].append(
|
||||
layer_loss_save)
|
||||
else:
|
||||
windows = [
|
||||
STA_param[head_idx + start_head] for head_idx in range(head_num)
|
||||
]
|
||||
|
||||
hidden_states = sliding_tile_attention(
|
||||
query, key, value, windows, text_length, has_text,
|
||||
self.dit_seq_shape_str).transpose(1, 2)
|
||||
windows = [
|
||||
self.mask_strategy[timestep][layer_idx][head_idx + start_head]
|
||||
for head_idx in range(head_num)
|
||||
]
|
||||
# if has_text is False:
|
||||
# from IPython import embed
|
||||
# embed()
|
||||
hidden_states = sliding_tile_attention(
|
||||
query, key, value, windows, text_length, has_text,
|
||||
self.img_latent_shape_str).transpose(1, 2)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -1,198 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
|
||||
try:
|
||||
from vsa import video_sparse_attn
|
||||
except ImportError:
|
||||
video_sparse_attn = None
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class VideoSparseAttentionBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> List[int]:
|
||||
return [64, 128]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "VIDEO_SPARSE_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> Type["VideoSparseAttentionImpl"]:
|
||||
return VideoSparseAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> Type["VideoSparseAttentionMetadata"]:
|
||||
return VideoSparseAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> Type["VideoSparseAttentionMetadataBuilder"]:
|
||||
return VideoSparseAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoSparseAttentionMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
dit_seq_shape: List[int]
|
||||
VSA_sparsity: float
|
||||
|
||||
|
||||
class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def prepare(self):
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
forward_batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VideoSparseAttentionMetadata:
|
||||
if forward_batch.latents is None:
|
||||
raise ValueError("latents cannot be None")
|
||||
|
||||
raw_latent_shape = forward_batch.raw_latent_shape
|
||||
if raw_latent_shape is None:
|
||||
raise ValueError("raw_latent_shape cannot be None")
|
||||
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.patch_size
|
||||
dit_seq_shape = [
|
||||
raw_latent_shape[2] // patch_size[0],
|
||||
raw_latent_shape[3] // patch_size[1],
|
||||
raw_latent_shape[4] // patch_size[2]
|
||||
]
|
||||
VSA_sparsity = forward_batch.VSA_sparsity
|
||||
|
||||
return VideoSparseAttentionMetadata(current_timestep=current_timestep,
|
||||
dit_seq_shape=dit_seq_shape,
|
||||
VSA_sparsity=VSA_sparsity)
|
||||
|
||||
|
||||
class VideoSparseAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.prefix = prefix
|
||||
sp_group = get_sp_group()
|
||||
self.sp_size = sp_group.world_size
|
||||
self.VSA_base_tile_size = [4, 4, 4]
|
||||
self.dit_seq_shape: List[int]
|
||||
self.full_window_size: List[int]
|
||||
self.img_seq_length: int
|
||||
|
||||
def tile(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = rearrange(x,
|
||||
"b (sp t h w) head d -> b (t sp h w) head d",
|
||||
sp=self.sp_size,
|
||||
t=self.dit_seq_shape[0] // self.sp_size,
|
||||
h=self.dit_seq_shape[1],
|
||||
w=self.dit_seq_shape[2])
|
||||
|
||||
return rearrange(
|
||||
x,
|
||||
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
|
||||
n_t=self.full_window_size[0],
|
||||
n_h=self.full_window_size[1],
|
||||
n_w=self.full_window_size[2],
|
||||
ts_t=self.VSA_base_tile_size[0],
|
||||
ts_h=self.VSA_base_tile_size[1],
|
||||
ts_w=self.VSA_base_tile_size[2])
|
||||
|
||||
def untile(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = rearrange(
|
||||
x,
|
||||
"b (n_t n_h n_w ts_t ts_h ts_w) h d -> b (n_t ts_t n_h ts_h n_w ts_w) h d",
|
||||
n_t=self.full_window_size[0],
|
||||
n_h=self.full_window_size[1],
|
||||
n_w=self.full_window_size[2],
|
||||
ts_t=self.VSA_base_tile_size[0],
|
||||
ts_h=self.VSA_base_tile_size[1],
|
||||
ts_w=self.VSA_base_tile_size[2])
|
||||
return rearrange(x,
|
||||
"b (t sp h w) head d -> b (sp t h w) head d",
|
||||
sp=self.sp_size,
|
||||
t=self.dit_seq_shape[0] // self.sp_size,
|
||||
h=self.dit_seq_shape[1],
|
||||
w=self.dit_seq_shape[2])
|
||||
|
||||
def preprocess_qkv(
|
||||
self,
|
||||
qkv: torch.Tensor,
|
||||
attn_metadata: VideoSparseAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
self.dit_seq_shape = attn_metadata.dit_seq_shape
|
||||
self.full_window_size = [
|
||||
self.dit_seq_shape[0] // self.VSA_base_tile_size[0],
|
||||
self.dit_seq_shape[1] // self.VSA_base_tile_size[1],
|
||||
self.dit_seq_shape[2] // self.VSA_base_tile_size[2]
|
||||
]
|
||||
self.img_seq_length = math.prod(self.dit_seq_shape)
|
||||
return self.tile(qkv)
|
||||
|
||||
def postprocess_output(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
attn_metadata: VideoSparseAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
return self.untile(output)
|
||||
|
||||
def forward( # type: ignore[override]
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
gate_compress: torch.Tensor,
|
||||
attn_metadata: VideoSparseAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
query = query.transpose(1, 2).contiguous()
|
||||
key = key.transpose(1, 2).contiguous()
|
||||
value = value.transpose(1, 2).contiguous()
|
||||
gate_compress = gate_compress.transpose(1, 2).contiguous()
|
||||
|
||||
VSA_sparsity = attn_metadata.VSA_sparsity
|
||||
|
||||
cur_topk = math.ceil(
|
||||
(1 - VSA_sparsity) *
|
||||
(self.img_seq_length / math.prod(self.VSA_base_tile_size)))
|
||||
|
||||
if video_sparse_attn is None:
|
||||
raise NotImplementedError("video_sparse_attn is not installed")
|
||||
|
||||
hidden_states = video_sparse_attn(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
topk=cur_topk,
|
||||
block_size=(4, 4, 4),
|
||||
compress_attn_weight=gate_compress).transpose(1, 2)
|
||||
|
||||
return hidden_states
|
||||
@@ -9,11 +9,10 @@ from fastvideo.v1.attention.selector import (backend_name_to_enum,
|
||||
get_attn_backend)
|
||||
from fastvideo.v1.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather, sequence_model_parallel_all_to_all_4D)
|
||||
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
|
||||
get_sp_world_size)
|
||||
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 AttentionBackendEnum
|
||||
from fastvideo.v1.utils import get_compute_dtype
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
class DistributedAttention(nn.Module):
|
||||
@@ -26,8 +25,8 @@ class DistributedAttention(nn.Module):
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[Tuple[
|
||||
AttentionBackendEnum, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
@@ -39,19 +38,19 @@ class DistributedAttention(nn.Module):
|
||||
if num_kv_heads is None:
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = get_compute_dtype()
|
||||
dtype = torch.get_default_dtype()
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
causal=causal,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
prefix=f"{prefix}.impl",
|
||||
**extra_impl_args)
|
||||
self.impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
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
|
||||
self.num_kv_heads = num_kv_heads
|
||||
@@ -86,8 +85,8 @@ class DistributedAttention(nn.Module):
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim(
|
||||
) == 4, "Expected 4D tensors"
|
||||
batch_size, seq_len, num_heads, head_dim = q.shape
|
||||
local_rank = get_sp_parallel_rank()
|
||||
world_size = get_sp_world_size()
|
||||
local_rank = get_sequence_model_parallel_rank()
|
||||
world_size = get_sequence_model_parallel_world_size()
|
||||
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
@@ -100,7 +99,7 @@ class DistributedAttention(nn.Module):
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
# Apply backend-specific preprocess_qkv
|
||||
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
|
||||
qkv = self.impl.preprocess_qkv(qkv, ctx_attn_metadata)
|
||||
|
||||
# Concatenate with replicated QKV if provided
|
||||
if replicated_q is not None:
|
||||
@@ -116,7 +115,7 @@ class DistributedAttention(nn.Module):
|
||||
|
||||
q, k, v = qkv.chunk(3, dim=0)
|
||||
|
||||
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||
output = self.impl.forward(q, k, v, ctx_attn_metadata)
|
||||
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
@@ -127,73 +126,7 @@ class DistributedAttention(nn.Module):
|
||||
replicated_output = sequence_model_parallel_all_gather(
|
||||
replicated_output.contiguous(), dim=2)
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output,
|
||||
scatter_dim=1,
|
||||
gather_dim=2)
|
||||
return output, replicated_output
|
||||
|
||||
|
||||
class DistributedAttention_VSA(DistributedAttention):
|
||||
"""Distributed attention layer with VSA support.
|
||||
"""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
replicated_q: Optional[torch.Tensor] = None,
|
||||
replicated_k: Optional[torch.Tensor] = None,
|
||||
replicated_v: Optional[torch.Tensor] = None,
|
||||
gate_compress: Optional[torch.Tensor] = None,
|
||||
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
"""Forward pass for distributed attention.
|
||||
|
||||
Args:
|
||||
q (torch.Tensor): Query tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
k (torch.Tensor): Key tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
v (torch.Tensor): Value tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
gate_compress (torch.Tensor): Gate compress tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
|
||||
replicated_k (Optional[torch.Tensor]): Replicated key tensor
|
||||
replicated_v (Optional[torch.Tensor]): Replicated value tensor
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
|
||||
- o (torch.Tensor): Output tensor after attention for the main sequence
|
||||
- replicated_o (Optional[torch.Tensor]): Output tensor for replicated tokens, if provided
|
||||
"""
|
||||
# Check text tokens are not supported for VSA now
|
||||
assert replicated_q is None and replicated_k is None and replicated_v is None, "Replicated QKV is not supported for VSA now"
|
||||
# Check input shapes
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim(
|
||||
) == 4, "Expected 4D tensors"
|
||||
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
# Stack QKV
|
||||
qkvg = torch.cat([q, k, v, gate_compress],
|
||||
dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
|
||||
# Redistribute heads across sequence dimension
|
||||
qkvg = sequence_model_parallel_all_to_all_4D(qkvg,
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
|
||||
qkvg = self.attn_impl.preprocess_qkv(qkvg, ctx_attn_metadata)
|
||||
|
||||
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
|
||||
output = self.attn_impl.forward(
|
||||
q, k, v, gate_compress, ctx_attn_metadata) # type: ignore[call-arg]
|
||||
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
|
||||
output = self.impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output,
|
||||
scatter_dim=1,
|
||||
@@ -211,8 +144,8 @@ class LocalAttention(nn.Module):
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[Tuple[
|
||||
AttentionBackendEnum, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
if softmax_scale is None:
|
||||
@@ -222,18 +155,18 @@ class LocalAttention(nn.Module):
|
||||
if num_kv_heads is None:
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = get_compute_dtype()
|
||||
dtype = torch.get_default_dtype()
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
causal=causal,
|
||||
**extra_impl_args)
|
||||
self.impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
causal=causal,
|
||||
**extra_impl_args)
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.num_kv_heads = num_kv_heads
|
||||
@@ -264,5 +197,5 @@ class LocalAttention(nn.Module):
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||
output = self.impl.forward(q, k, v, ctx_attn_metadata)
|
||||
return output
|
||||
|
||||
@@ -11,13 +11,13 @@ import torch
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.attention.backends.abstract import AttentionBackend
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum, current_platform
|
||||
from fastvideo.v1.platforms import _Backend, current_platform
|
||||
from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def backend_name_to_enum(backend_name: str) -> Optional[AttentionBackendEnum]:
|
||||
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
|
||||
"""
|
||||
Convert a string backend name to a _Backend enum value.
|
||||
|
||||
@@ -27,11 +27,11 @@ def backend_name_to_enum(backend_name: str) -> Optional[AttentionBackendEnum]:
|
||||
loaded.
|
||||
"""
|
||||
assert backend_name is not None
|
||||
return AttentionBackendEnum[backend_name] if backend_name in AttentionBackendEnum.__members__ else \
|
||||
return _Backend[backend_name] if backend_name in _Backend.__members__ else \
|
||||
None
|
||||
|
||||
|
||||
def get_env_variable_attn_backend() -> Optional[AttentionBackendEnum]:
|
||||
def get_env_variable_attn_backend() -> Optional[_Backend]:
|
||||
'''
|
||||
Get the backend override specified by the FastVideo attention
|
||||
backend environment variable, if one is specified.
|
||||
@@ -53,11 +53,10 @@ def get_env_variable_attn_backend() -> Optional[AttentionBackendEnum]:
|
||||
#
|
||||
# THIS SELECTION TAKES PRECEDENCE OVER THE
|
||||
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
|
||||
forced_attn_backend: Optional[AttentionBackendEnum] = None
|
||||
forced_attn_backend: Optional[_Backend] = None
|
||||
|
||||
|
||||
def global_force_attn_backend(
|
||||
attn_backend: Optional[AttentionBackendEnum]) -> None:
|
||||
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
|
||||
'''
|
||||
Force all attention operations to use a specified backend.
|
||||
|
||||
@@ -72,7 +71,7 @@ def global_force_attn_backend(
|
||||
forced_attn_backend = attn_backend
|
||||
|
||||
|
||||
def get_global_forced_attn_backend() -> Optional[AttentionBackendEnum]:
|
||||
def get_global_forced_attn_backend() -> Optional[_Backend]:
|
||||
'''
|
||||
Get the currently-forced choice of attention backend,
|
||||
or None if auto-selection is currently enabled.
|
||||
@@ -83,8 +82,7 @@ def get_global_forced_attn_backend() -> Optional[AttentionBackendEnum]:
|
||||
def get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
return _cached_get_attn_backend(head_size, dtype,
|
||||
supported_attention_backends)
|
||||
@@ -94,8 +92,7 @@ def get_attn_backend(
|
||||
def _cached_get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
@@ -105,7 +102,7 @@ def _cached_get_attn_backend(
|
||||
if not supported_attention_backends:
|
||||
raise ValueError("supported_attention_backends is empty")
|
||||
selected_backend = None
|
||||
backend_by_global_setting: Optional[AttentionBackendEnum] = (
|
||||
backend_by_global_setting: Optional[_Backend] = (
|
||||
get_global_forced_attn_backend())
|
||||
if backend_by_global_setting is not None:
|
||||
selected_backend = backend_by_global_setting
|
||||
@@ -128,7 +125,7 @@ def _cached_get_attn_backend(
|
||||
|
||||
@contextmanager
|
||||
def global_force_attn_backend_context_manager(
|
||||
attn_backend: AttentionBackendEnum) -> Generator[None, None, None]:
|
||||
attn_backend: _Backend) -> Generator[None, None, None]:
|
||||
'''
|
||||
Globally force a FastVideo attention backend override within a
|
||||
context manager, reverting the global attention backend
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import Any, Dict
|
||||
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, List, Optional, Tuple
|
||||
from typing import Any, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -12,16 +11,15 @@ class DiTArchConfig(ArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=list)
|
||||
_compile_conditions: list = field(default_factory=list)
|
||||
_param_names_mapping: dict = field(default_factory=dict)
|
||||
_lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN)
|
||||
_supported_attention_backends: Tuple[_Backend,
|
||||
...] = (_Backend.SLIDING_TILE_ATTN,
|
||||
_Backend.SAGE_ATTN,
|
||||
_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 0
|
||||
exclude_lora_layers: List[str] = field(default_factory=list)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self._compile_conditions:
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional, Tuple
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -164,8 +163,6 @@ class HunyuanVideoArchConfig(DiTArchConfig):
|
||||
pooled_projection_dim: int = 768
|
||||
rope_theta: int = 256
|
||||
qk_norm: str = "rms_norm"
|
||||
exclude_lora_layers: List[str] = field(
|
||||
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
@@ -52,7 +51,6 @@ class StepVideoArchConfig(DiTArchConfig):
|
||||
default_factory=lambda: [6144, 1024])
|
||||
attention_type: Optional[str] = "torch"
|
||||
use_additional_conditions: Optional[bool] = False
|
||||
exclude_lora_layers: List[str] = field(default_factory=lambda: [])
|
||||
|
||||
def __post_init__(self):
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional, Tuple
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
@@ -52,23 +51,6 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
r"blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
})
|
||||
# Some LoRA adapters use the original official layer names instead of hf layer names,
|
||||
# so apply this before the param_names_mapping
|
||||
_lora_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$":
|
||||
r"blocks.\1.attn1.to_out.0.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$":
|
||||
r"blocks.\1.attn2.to_out.0.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
|
||||
})
|
||||
|
||||
patch_size: Tuple[int, int, int] = (1, 2, 2)
|
||||
text_len = 512
|
||||
@@ -86,7 +68,6 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
image_dim: Optional[int] = None
|
||||
added_kv_proj_dim: Optional[int] = None
|
||||
rope_max_seq_len: int = 1024
|
||||
exclude_lora_layers: List[str] = field(default_factory=lambda: ["embedder"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
@@ -6,14 +5,14 @@ import torch
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
@dataclass
|
||||
class EncoderArchConfig(ArchConfig):
|
||||
architectures: List[str] = field(default_factory=lambda: [])
|
||||
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
|
||||
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA)
|
||||
output_hidden_states: bool = False
|
||||
use_return_dict: bool = True
|
||||
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
|
||||
@@ -1,6 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import dataclasses
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Union
|
||||
|
||||
@@ -131,12 +128,3 @@ class VAEConfig(ModelConfig):
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "VAEConfig":
|
||||
kwargs = {}
|
||||
for attr in dataclasses.fields(cls):
|
||||
value = getattr(args, attr.name, None)
|
||||
if value is not None:
|
||||
kwargs[attr.name] = value
|
||||
return cls(**kwargs)
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Tuple
|
||||
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Tuple
|
||||
|
||||
@@ -64,7 +63,7 @@ class WanVAEArchConfig(VAEArchConfig):
|
||||
|
||||
@dataclass
|
||||
class WanVAEConfig(VAEConfig):
|
||||
arch_config: WanVAEArchConfig = field(default_factory=WanVAEArchConfig)
|
||||
arch_config: VAEArchConfig = field(default_factory=WanVAEArchConfig)
|
||||
use_feature_cache: bool = True
|
||||
|
||||
use_tiling: bool = False
|
||||
|
||||
@@ -3,7 +3,7 @@ from fastvideo.v1.configs.pipelines.base import (PipelineConfig,
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
|
||||
HunyuanConfig)
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
get_pipeline_config_cls_for_name)
|
||||
from fastvideo.v1.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
|
||||
WanI2V720PConfig,
|
||||
@@ -14,5 +14,5 @@ __all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
"get_pipeline_config_cls_for_name"
|
||||
]
|
||||
|
||||
@@ -1,17 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
from dataclasses import asdict, dataclass, field, fields
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union, cast
|
||||
from typing import Any, Callable, Dict, Optional, Tuple, cast
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
VAEConfig)
|
||||
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.v1.configs.utils import update_config_from_args
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import (FlexibleArgumentParser, StoreBoolean,
|
||||
shallow_asdict)
|
||||
from fastvideo.v1.utils import shallow_asdict
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -24,282 +21,58 @@ def postprocess_text(output: BaseEncoderOutput) -> torch.tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
# config for a single pipeline
|
||||
@dataclass
|
||||
class PipelineConfig:
|
||||
"""Base configuration for all pipeline architectures."""
|
||||
model_path: str = ""
|
||||
pipeline_config_path: Optional[str] = None
|
||||
|
||||
# Video generation parameters
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
use_cpu_offload: bool = False
|
||||
disable_autocast: bool = False
|
||||
|
||||
# Model configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
dit_precision: str = "bf16"
|
||||
precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = True
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
|
||||
image_encoder_precision: str = "fp32"
|
||||
# DiT configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
|
||||
# Text encoder configuration
|
||||
DEFAULT_TEXT_ENCODER_PRECISIONS = ("fp16", )
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (EncoderConfig(), ))
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", ))
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (EncoderConfig(), ))
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(postprocess_text, ))
|
||||
|
||||
# LoRA parameters
|
||||
lora_path: Optional[str] = None
|
||||
lora_nickname: Optional[
|
||||
str] = "default" # for swapping adapters in the pipeline
|
||||
lora_target_names: Optional[List[
|
||||
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
|
||||
|
||||
# StepVideo specific parameters
|
||||
pos_magic: Optional[str] = None
|
||||
neg_magic: Optional[str] = None
|
||||
timesteps_scale: Optional[bool] = None
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
STA_mode: Optional[str] = None
|
||||
skip_time_steps: int = 15
|
||||
|
||||
# Compilation
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser,
|
||||
prefix: str = "") -> FlexibleArgumentParser:
|
||||
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
|
||||
|
||||
# model_path will be conflicting with the model_path in FastVideoArgs,
|
||||
# so we add it separately if prefix is not empty
|
||||
if prefix_with_dot != "":
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}model-path",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}model_path",
|
||||
default=PipelineConfig.model_path,
|
||||
help="Path to the pretrained model",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}pipeline-config-path",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}pipeline_config_path",
|
||||
default=PipelineConfig.pipeline_config_path,
|
||||
help="Path to the pipeline config",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}embedded-cfg-scale",
|
||||
type=float,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}embedded_cfg_scale",
|
||||
default=PipelineConfig.embedded_cfg_scale,
|
||||
help="Embedded CFG scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}flow-shift",
|
||||
type=float,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}flow_shift",
|
||||
default=PipelineConfig.flow_shift,
|
||||
help="Flow shift parameter",
|
||||
)
|
||||
|
||||
# DiT configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}dit-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}dit_precision",
|
||||
default=PipelineConfig.dit_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for the DiT model",
|
||||
)
|
||||
|
||||
# VAE configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_precision",
|
||||
default=PipelineConfig.vae_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for VAE",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-tiling",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_tiling",
|
||||
default=PipelineConfig.vae_tiling,
|
||||
help="Enable VAE tiling",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-sp",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_sp",
|
||||
help="Enable VAE spatial parallelism",
|
||||
)
|
||||
|
||||
# Text encoder configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}text-encoder-precisions",
|
||||
nargs="+",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}text_encoder_precisions",
|
||||
default=PipelineConfig.DEFAULT_TEXT_ENCODER_PRECISIONS,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for each text encoder",
|
||||
)
|
||||
|
||||
# Image encoder configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}image-encoder-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}image_encoder_precision",
|
||||
default=PipelineConfig.image_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for image encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}pos_magic",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}pos_magic",
|
||||
default=PipelineConfig.pos_magic,
|
||||
help="Positive magic prompt for sampling, used in stepvideo",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}neg_magic",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}neg_magic",
|
||||
default=PipelineConfig.neg_magic,
|
||||
help="Negative magic prompt for sampling, used in stepvideo",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}timesteps_scale",
|
||||
type=bool,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}timesteps_scale",
|
||||
default=PipelineConfig.timesteps_scale,
|
||||
help=
|
||||
"Bool for applying scheduler scale in set_timesteps, used in stepvideo",
|
||||
)
|
||||
|
||||
# Add VAE configuration arguments
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig
|
||||
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
|
||||
|
||||
# Add DiT configuration arguments
|
||||
from fastvideo.v1.configs.models.dits.base import DiTConfig
|
||||
DiTConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}dit-config")
|
||||
|
||||
return parser
|
||||
|
||||
def update_config_from_dict(self,
|
||||
args: Dict[str, Any],
|
||||
prefix: str = "") -> None:
|
||||
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
|
||||
update_config_from_args(self, args, prefix, pop_args=True)
|
||||
update_config_from_args(self.vae_config,
|
||||
args,
|
||||
f"{prefix_with_dot}vae_config",
|
||||
pop_args=True)
|
||||
update_config_from_args(self.dit_config,
|
||||
args,
|
||||
f"{prefix_with_dot}dit_config",
|
||||
pop_args=True)
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "PipelineConfig":
|
||||
"""
|
||||
use the pipeline class setting from model_path to match the pipeline config
|
||||
"""
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
|
||||
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
|
||||
|
||||
@classmethod
|
||||
def from_kwargs(cls,
|
||||
kwargs: Dict[str, Any],
|
||||
config_cli_prefix: str = "") -> "PipelineConfig":
|
||||
"""
|
||||
Load PipelineConfig from kwargs Dictionary.
|
||||
kwargs: dictionary of kwargs
|
||||
config_cli_prefix: prefix of CLI arguments for this PipelineConfig instance
|
||||
"""
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
|
||||
prefix_with_dot = f"{config_cli_prefix}." if (config_cli_prefix.strip()
|
||||
!= "") else ""
|
||||
model_path: Optional[str] = kwargs.get(prefix_with_dot + 'model_path',
|
||||
None) or kwargs.get('model_path')
|
||||
pipeline_config_or_path: Optional[Union[str, PipelineConfig, Dict[
|
||||
str, Any]]] = kwargs.get(prefix_with_dot + 'pipeline_config',
|
||||
None) or kwargs.get('pipeline_config')
|
||||
if model_path is None:
|
||||
raise ValueError("model_path is required in kwargs")
|
||||
|
||||
# 1. Get the pipeline config class from the registry
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
|
||||
# 2. Instantiate PipelineConfig
|
||||
if pipeline_config_cls is None:
|
||||
get_pipeline_config_cls_for_name)
|
||||
pipeline_config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if pipeline_config_cls is not None:
|
||||
pipeline_config = pipeline_config_cls()
|
||||
else:
|
||||
logger.warning(
|
||||
"Couldn't find pipeline config for %s. Using the default pipeline config.",
|
||||
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
|
||||
model_path)
|
||||
pipeline_config = cls()
|
||||
else:
|
||||
pipeline_config = pipeline_config_cls()
|
||||
|
||||
# 3. Load PipelineConfig from a json file or a PipelineConfig object if provided
|
||||
if isinstance(pipeline_config_or_path, str):
|
||||
pipeline_config.load_from_json(pipeline_config_or_path)
|
||||
kwargs[prefix_with_dot +
|
||||
'pipeline_config_path'] = pipeline_config_or_path
|
||||
elif isinstance(pipeline_config_or_path, PipelineConfig):
|
||||
pipeline_config = pipeline_config_or_path
|
||||
elif isinstance(pipeline_config_or_path, dict):
|
||||
pipeline_config.update_pipeline_config(pipeline_config_or_path)
|
||||
|
||||
# 4. Update PipelineConfig from CLI arguments if provided
|
||||
kwargs[prefix_with_dot + 'model_path'] = model_path
|
||||
pipeline_config.update_config_from_dict(kwargs, config_cli_prefix)
|
||||
return pipeline_config
|
||||
|
||||
def check_pipeline_config(self) -> None:
|
||||
if self.vae_sp and not self.vae_tiling:
|
||||
raise ValueError(
|
||||
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
return cast(PipelineConfig, pipeline_config)
|
||||
|
||||
def dump_to_json(self, file_path: str):
|
||||
output_dict = shallow_asdict(self)
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Tuple, TypedDict
|
||||
|
||||
@@ -69,6 +68,9 @@ class HunyuanConfig(PipelineConfig):
|
||||
embedded_cfg_scale: int = 6
|
||||
flow_shift: int = 7
|
||||
|
||||
# Video parameters
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Text encoding stage
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (LlamaConfig(), CLIPTextConfig()))
|
||||
@@ -80,7 +82,7 @@ class HunyuanConfig(PipelineConfig):
|
||||
(llama_postprocess_text, clip_postprocess_text))
|
||||
|
||||
# Precision for each component
|
||||
dit_precision: str = "bf16"
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", "fp16"))
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Registry for pipeline weight-specific configurations."""
|
||||
|
||||
import os
|
||||
@@ -19,7 +18,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Registry maps specific model weights to their config classes
|
||||
PIPE_NAME_TO_CONFIG: Dict[str, Type[PipelineConfig]] = {
|
||||
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
|
||||
@@ -51,74 +50,37 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
|
||||
}
|
||||
|
||||
|
||||
def get_pipeline_config_cls_from_name(
|
||||
pipeline_name_or_path: str) -> Type[PipelineConfig]:
|
||||
"""Get the appropriate configuration class for a given pipeline name or path.
|
||||
def get_pipeline_config_cls_for_name(
|
||||
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
|
||||
"""Get the appropriate config class for specific pretrained weights."""
|
||||
|
||||
This function implements a multi-step lookup process to find the most suitable
|
||||
configuration class for a given pipeline. It follows this order:
|
||||
1. Exact match in the PIPE_NAME_TO_CONFIG
|
||||
2. Partial match in the PIPE_NAME_TO_CONFIG
|
||||
3. Fallback to class name in the model_index.json
|
||||
4. else raise an error
|
||||
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)
|
||||
|
||||
Args:
|
||||
pipeline_name_or_path (str): The name or path of the pipeline. This can be:
|
||||
- A registered model ID (e.g., "FastVideo/FastHunyuan-diffusers")
|
||||
- A local path to a model directory
|
||||
- A model ID that will be downloaded
|
||||
|
||||
Returns:
|
||||
Type[PipelineConfig]: The configuration class that best matches the pipeline.
|
||||
This will be one of:
|
||||
- A specific weight configuration class if an exact match is found
|
||||
- A fallback configuration class based on the pipeline architecture
|
||||
- The base PipelineConfig class if no matches are found
|
||||
|
||||
Note:
|
||||
- For local paths, the function will verify the model configuration
|
||||
- For remote models, it will attempt to download the model index
|
||||
- Warning messages are logged when falling back to less specific configurations
|
||||
"""
|
||||
|
||||
pipeline_config_cls: Optional[Type[PipelineConfig]] = None
|
||||
pipeline_name = config["_class_name"]
|
||||
|
||||
# First try exact match for specific weights
|
||||
if pipeline_name_or_path in PIPE_NAME_TO_CONFIG:
|
||||
pipeline_config_cls = PIPE_NAME_TO_CONFIG[pipeline_name_or_path]
|
||||
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 PIPE_NAME_TO_CONFIG.items():
|
||||
for registered_id, config_class in WEIGHT_CONFIG_REGISTRY.items():
|
||||
if registered_id in pipeline_name_or_path:
|
||||
pipeline_config_cls = config_class
|
||||
break
|
||||
return config_class
|
||||
|
||||
# If no match, try to use the fallback config
|
||||
if pipeline_config_cls is None:
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
config = verify_model_config_and_directory(pipeline_name_or_path)
|
||||
else:
|
||||
config = maybe_download_model_index(pipeline_name_or_path)
|
||||
logger.warning(
|
||||
"Trying to use the config from the model_index.json. FastVideo may not correctly identify the optimal config for this model in this situation."
|
||||
)
|
||||
fallback_config = None
|
||||
# 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
|
||||
|
||||
pipeline_name = config["_class_name"]
|
||||
# Try to determine pipeline architecture for fallback
|
||||
for pipeline_type, detector in PIPELINE_DETECTOR.items():
|
||||
if detector(pipeline_name.lower()):
|
||||
pipeline_config_cls = PIPELINE_FALLBACK_CONFIG.get(
|
||||
pipeline_type)
|
||||
break
|
||||
|
||||
if pipeline_config_cls is not None:
|
||||
logger.warning(
|
||||
"No match found for pipeline %s, using fallback config %s.",
|
||||
pipeline_name_or_path, pipeline_config_cls)
|
||||
|
||||
if pipeline_config_cls is None:
|
||||
raise ValueError(
|
||||
f"No match found for pipeline {pipeline_name_or_path}, please check the pipeline name or path."
|
||||
)
|
||||
|
||||
return pipeline_config_cls
|
||||
logger.warning("No match found for pipeline %s, using fallback config %s.",
|
||||
pipeline_name_or_path, fallback_config)
|
||||
return fallback_config
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, VAEConfig
|
||||
@@ -19,6 +18,9 @@ class StepVideoT2VConfig(PipelineConfig):
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# Video parameters
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 13
|
||||
timesteps_scale: bool = False
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Tuple
|
||||
|
||||
@@ -38,6 +37,9 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# Video parameters
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 3
|
||||
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user