Compare commits

..
Author SHA1 Message Date
rlsu9 805978d0da stage bash diff 2025-03-13 17:11:55 +00:00
rlsu9 797cb2f659 recover sample hunyuan 2025-03-13 17:08:02 +00:00
rlsu9 ce749878be fix attention gap 2025-03-13 17:04:53 +00:00
rlsu9 8c844ba8ac stage change 2025-03-13 16:42:43 +00:00
rlsu9 69acec287e stage sh 2025-03-13 16:19:44 +00:00
rlsu9 5a2fd28f95 fix cache lime error 2025-03-13 08:15:18 +00:00
rlsu9 3133dbf69d remove unused item 2025-03-13 07:25:47 +00:00
rlsu9 0d3a05d58e remove unused difference 2025-03-13 07:16:39 +00:00
rlsu9 f1fc13b28d remove unused files 2025-03-13 07:05:39 +00:00
rlsu9 0b4033decc upload debug code 2025-03-13 07:02:11 +00:00
rlsu9 6b119a9f76 stage 2025-03-13 05:38:29 +00:00
rlsu9 15fb1d9a57 add tk debug 2025-03-07 01:29:22 +00:00
rlsu9 66060dd22f update kernel benchmark 2025-03-06 23:59:00 +00:00
rlsu9 9fe8704537 stage kernel-flops benchmark 2025-03-06 23:23:50 +00:00
rlsu9 31bbdcd2a4 stage 2025-03-04 00:51:21 +00:00
96 changed files with 96 additions and 18647 deletions
-70
View File
@@ -1,70 +0,0 @@
name: Publish FastVideo to PyPI on Version Change
on:
push:
branches:
- main
paths:
- 'pyproject.toml' # Trigger when pyproject.toml changes
jobs:
check-version-change:
runs-on: ubuntu-latest
outputs:
version-changed: ${{ steps.check-version.outputs.changed }}
new-version: ${{ steps.check-version.outputs.new-version }}
steps:
- name: Checkout code
uses: actions/checkout@v3
with:
fetch-depth: 2
- name: Check if version changed
id: check-version
run: |
# Get current commit's version
NEW_VERSION=$(grep -oP 'version\s*=\s*"\K[^"]+' pyproject.toml)
echo "New version: $NEW_VERSION"
# Get previous version from git history
OLD_VERSION=$(git show HEAD~1:./pyproject.toml | grep -oP 'version\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
echo "changed=true" >> $GITHUB_OUTPUT
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
else
echo "Version did not change"
echo "changed=false" >> $GITHUB_OUTPUT
fi
build-publish-main:
needs: check-version-change
if: needs.check-version-change.outputs.version-changed == 'true'
runs-on: ubuntu-latest
permissions:
id-token: write # Needed for OIDC Trusted Publishing
steps:
- name: Checkout code
uses: actions/checkout@v3
- name: Set up Python
uses: actions/setup-python@v4
with:
python-version: '3.10'
- name: Install build dependencies
run: |
python -m pip install --upgrade pip
pip install build twine wheel
- name: Build package
run: |
python -m build
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: dist/
-221
View File
@@ -1,221 +0,0 @@
name: Publish Sliding Tile Attention Kernel to PyPI on Version Change
on:
push:
branches:
- main
paths:
- "csrc/sliding_tile_attention/setup.py"
jobs:
check-version-change:
runs-on: ubuntu-latest
outputs:
version-changed: ${{ steps.check-version.outputs.changed }}
new-version: ${{ steps.check-version.outputs.new-version }}
steps:
- name: Checkout code
uses: actions/checkout@v3
with:
fetch-depth: 2
- name: Check if version changed
id: check-version
run: |
cd csrc/sliding_tile_attention
# Get current commit's version
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.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
echo "changed=true" >> $GITHUB_OUTPUT
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
else
echo "Version did not change"
echo "changed=false" >> $GITHUB_OUTPUT
fi
build_wheels:
name: Build Wheel
needs: check-version-change
if: needs.check-version-change.outputs.version-changed == 'true'
runs-on: ${{ matrix.os }}
strategy:
fail-fast: false
matrix:
# Using ubuntu-20.04 instead of 22.04 for more compatibility (glibc). Ideally we'd use the
# manylinux docker image, but I haven't figured out how to install CUDA on manylinux.
os: [ubuntu-22.04]
python-version: ['3.10', '3.11', '3.12', '3.13']
torch-version: ['2.5.1', '2.6.0']
cuda-version: ['12.4.1', '12.5.1', '12.6.3']
steps:
- name: Checkout
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
- name: Install CUDA ${{ matrix.cuda-version }}
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: ${{ matrix.cuda-version }}
linux-local-args: '["--toolkit"]'
method: 'network'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-${{ matrix.cuda-version }}
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch ${{ matrix.torch-version }}+cu${{ matrix.cuda-version }}
run: |
pip install --upgrade pip
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
pip install typing-extensions==4.12.2
# We want to figure out the CUDA version to download pytorch
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
export TORCH_CUDA_VERSION=124
pip install --no-cache-dir torch==${{ matrix.torch-version }} --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build wheel
run: |
# 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/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/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)
# Get the correct version format
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
# Rename with version information
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
- name: Upload wheel artifact
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}
path: csrc/sliding_tile_attention/dist/*.whl
retention-days: 90
publish_package:
name: Publish package
needs: [build_wheels]
if: needs.check-version-change.outputs.version-changed == 'true'
runs-on: ubuntu-22.04
permissions:
id-token: write # Needed for OIDC Trusted Publishing
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: '3.10'
- name: Install CUDA 12.4.1
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: 12.4.1
linux-local-args: '["--toolkit"]'
method: 'network'
sub-packages: '["nvcc"]'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-12.4.1
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch 2.5.1+cu12.4.1
run: |
pip install --upgrade pip
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
pip install typing-extensions==4.12.2
# We want to figure out the CUDA version to download pytorch
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
export TORCH_CUDA_VERSION=124
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build source distribution
run: |
# 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/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/sliding_tile_attention/dist/
-3
View File
@@ -23,11 +23,8 @@ jobs:
python -m pip install --upgrade pip setuptools wheel
pip install torch
pip install packaging ninja
# remove st-attn dependency because no cuda environment
sed -i '/st_attn/d' pyproject.toml
pip install -e .
pip install pytest
- name: Run Pytest
run: |
pytest --ignore csrc/sliding_tile_attention/test
-2
View File
@@ -1,2 +0,0 @@
recursive-include tk *
include config.py
+3 -20
View File
@@ -7,13 +7,6 @@ from torch.utils.cpp_extension import BuildExtension, CUDAExtension
target = target.lower()
# Package metadata
PACKAGE_NAME = "st_attn"
VERSION = "0.0.2"
AUTHOR = "Hao AI Lab"
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
# 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',
@@ -51,11 +44,8 @@ for k in kernels:
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,
setup(name='st_attn',
version="0.0.0",
packages=find_packages(),
ext_modules=[
CUDAExtension('st_attn_cuda',
@@ -66,11 +56,4 @@ setup(name=PACKAGE_NAME,
},
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"])
cmdclass={'build_ext': BuildExtension})
-2
View File
@@ -8,7 +8,5 @@ pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-is
pip install -r requirements-lint.txt
pip install -r requirements.txt
# install fastvideo
pip install -e .
+1 -1
View File
@@ -160,7 +160,7 @@ def distill_one_step(
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)
teacher_output = cond_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
@@ -813,7 +813,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
t, l, h = map(int, key.split('_'))
result[t][l][h] = value
return result
torch._dynamo.config.cache_size_limit = 128
mask_strategy = dict_to_3d_list(mask_strategy)
# if is_progress_bar:
with self.progress_bar(total=num_inference_steps) as progress_bar:
+85 -4
View File
@@ -6,12 +6,68 @@ try:
from st_attn import sliding_tile_attention
except ImportError:
print("Could not load Sliding Tile Attention.")
sliding_tile_attention = None
sliding_tile_attention = None
from functools import lru_cache
from torch.nn.attention.flex_attention import flex_attention
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
from csrc.sliding_tile_attention.test.flex_sta_ref import get_sliding_tile_attention_mask
@lru_cache(maxsize=32)
def get_compiled_flex_attention(strategy, tile_size, image_size, text_length, device):
"""
Create and compile flex attention with a specific sliding block mask.
This function is cached to avoid recompiling for the same parameters.
Args:
strategy (tuple): A tuple (t, h, w) defining the strategy
tile_size (tuple): A tuple (ts_t, ts_h, ts_w) defining the tile size
image_size (tuple): A tuple (n_t, n_h, n_w) defining the image size
text_length (int): The text length
device (str): The device to use
Returns:
function: A compiled flex attention function with the specified mask
"""
# Convert strategy to the required format (ceil(t*3/2), h*2, w)
adjusted_strategy = strategy
# Get the sliding block attention mask
mask = get_sliding_tile_attention_mask(
adjusted_strategy,
tile_size,
image_size,
text_length,
device
)
def flex_attn_with_mask(q, k, v, scale=None):
return flex_attention(q, k, v, block_mask=mask, scale=scale)
# Compile the wrapper function
compiled_flex_attn = torch.compile(flex_attn_with_mask)
return compiled_flex_attn
def flex_sliding_tile_attention(q_all, k_all, v_all, strategy, tile_size,
image_size, text_length, scale=None):
device = q_all.device
# Get the compiled flex attention function (cached if called with same parameters)
compiled_flex_attn = get_compiled_flex_attention(
strategy,
tile_size,
image_size,
text_length,
device
)
# Apply the compiled flex attention
output = compiled_flex_attn(q_all, k_all, v_all, scale=scale)
return output
def attention(
q,
@@ -92,7 +148,32 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
start_head = current_rank * head_num
windows = [mask_strategy[head_idx + start_head] for head_idx in range(head_num)]
hidden_states = sliding_tile_attention(query, key, value, windows, text_length).transpose(1, 2)
if sliding_tile_attention is not None:
hidden_states = sliding_tile_attention(query, key, value, windows, text_length).transpose(1, 2)
else:
print("Sliding Tile Attention not available. Using Flex Sliding Tile Attention.")
hidden_states = torch.empty_like(query)
strategy_to_heads = {}
for head_index in range(head_num):
strategy = tuple(windows[head_index]) # Convert list to tuple for dict key
if strategy not in strategy_to_heads:
strategy_to_heads[strategy] = []
strategy_to_heads[strategy].append(head_index)
for strategy, heads in strategy_to_heads.items():
# Gather all heads with this strategy
query_heads = torch.cat([query[:, head_idx:head_idx + 1, :, :] for head_idx in heads], dim=1)
key_heads = torch.cat([key[:, head_idx:head_idx + 1, :, :] for head_idx in heads], dim=1)
value_heads = torch.cat([value[:, head_idx:head_idx + 1, :, :] for head_idx in heads], dim=1)
# Process all heads with this strategy at once
# processed_heads = selected_attn_processor[processor_idx](query_heads, key_heads, value_heads)
processed_heads = flex_sliding_tile_attention(query_heads, key_heads, value_heads, strategy, (6, 8, 8), (30, 48, 80), text_length)
# Distribute results back to the correct positions
for i, head_idx in enumerate(heads):
hidden_states[:, head_idx:head_idx + 1, :, :] = processed_heads[:, i:i + 1, :, :]
hidden_states = hidden_states.transpose(1, 2)
else:
query = torch.cat([query, encoder_query], dim=1)
key = torch.cat([key, encoder_key], dim=1)
@@ -121,4 +202,4 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
b, s, a, d = attn.shape
attn = attn.reshape(b, s, -1)
return attn
return attn
@@ -523,8 +523,6 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if guidance is None:
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=torch.bfloat16)
if mask_strategy is None:
mask_strategy = [[None] * self.heads_num for _ in range(len(self.double_blocks) + len(self.single_blocks))]
img = x = hidden_states
text_mask = encoder_attention_mask
t = timestep
+1 -1
View File
@@ -207,4 +207,4 @@ if __name__ == "__main__":
# process for vae sequence parallel
if args.vae_sp and not args.vae_tiling:
raise ValueError("Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True.")
main(args)
main(args)
-3
View File
@@ -1,3 +0,0 @@
from .flash_attn import (DistributedAttention, LocalAttention)
__all__ = ["DistributedAttention", "LocalAttention"]
-138
View File
@@ -1,138 +0,0 @@
from itertools import accumulate
from typing import List, Optional
import torch
import torch.nn as nn
from fastvideo.v1.distributed.communication_op import sequence_model_parallel_all_to_all_4D, sequence_model_parallel_all_gather
from fastvideo.v1.distributed.parallel_state import get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size
from flash_attn import flash_attn_func, flash_attn_varlen_func
class DistributedAttention(nn.Module):
"""Distributed attention module that supports sequence parallelism.
This class implements a minimal attention operation with support for distributed
processing across multiple GPUs using sequence parallelism. The implementation assumes
batch_size=1 and no padding tokens for simplicity.
The sequence parallelism strategy follows the Ulysses paper (https://arxiv.org/abs/2309.14509),
which proposes redistributing attention heads across sequence dimension to enable efficient
parallel processing of long sequences.
Args:
dropout_rate (float, optional): Dropout probability. Defaults to 0.0.
causal (bool, optional): Whether to use causal attention. Defaults to False.
softmax_scale (float, optional): Custom scaling factor for attention scores.
If None, uses 1/sqrt(head_dim). Defaults to None.
"""
def __init__(
self,
dropout_rate: float = 0.0,
causal: bool = False,
softmax_scale: Optional[float] = None,
):
super().__init__()
self.dropout_rate = dropout_rate
self.causal = causal
self.softmax_scale = softmax_scale
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,
) -> 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]
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 input shapes
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
# assert bs = 1
assert q.shape[0] == 1, "Batch size must be 1, and there should be no padding tokens"
batch_size, seq_len, num_heads, head_dim = q.shape
local_rank = get_sequence_model_parallel_rank()
world_size = get_sequence_model_parallel_world_size()
# Stack QKV
qkv = torch.cat([q, k, v], dim=0) # [3, seq_len, num_heads, head_dim]
# Redistribute heads across sequence dimension
qkv = sequence_model_parallel_all_to_all_4D(qkv, scatter_dim=2, gather_dim=1)
# Concatenate with replicated QKV if provided
if replicated_q is not None:
assert replicated_k is not None and replicated_v is not None
replicated_qkv = torch.cat([replicated_q, replicated_k, replicated_v], dim=0) # [3, seq_len, num_heads, head_dim]
heads_per_rank = num_heads // world_size
replicated_qkv = replicated_qkv[:, :, local_rank * heads_per_rank:(local_rank + 1) * heads_per_rank]
qkv = torch.cat([qkv, replicated_qkv], dim=1)
q, k, v = qkv.chunk(3, dim=0)
# Apply flash attention
output = flash_attn_func(
q,
k,
v,
dropout_p=self.dropout_rate,
softmax_scale=self.softmax_scale,
causal=self.causal
)
# Redistribute back if using sequence parallelism
replicated_output = None
if replicated_q is not None:
replicated_output = output[:, seq_len*world_size:]
output = output[:, :seq_len*world_size]
# TODO: make this asynchronous
replicated_output = sequence_model_parallel_all_gather(replicated_output, dim=2)
output = sequence_model_parallel_all_to_all_4D(output, scatter_dim=1, gather_dim=2)
return output, replicated_output
class LocalAttention(nn.Module):
def __init__(self, dropout_rate: float = 0.0, causal: bool = False, softmax_scale: Optional[float] = None):
super().__init__()
self.dropout_rate = dropout_rate
self.causal = causal
self.softmax_scale = softmax_scale
def forward(self, q, k, v):
"""
Apply local attention between query, key and value tensors.
Args:
q (torch.Tensor): Query tensor of shape [batch_size, seq_len, num_heads, head_dim]
k (torch.Tensor): Key tensor of shape [batch_size, seq_len, num_heads, head_dim]
v (torch.Tensor): Value tensor of shape [batch_size, seq_len, num_heads, head_dim]
Returns:
torch.Tensor: Output tensor after local attention
"""
# Check input shapes
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
# Apply flash attention
output = flash_attn_func(
q,
k,
v,
dropout_p=self.dropout_rate,
softmax_scale=self.softmax_scale,
causal=self.causal
)
return output
-5
View File
@@ -1,5 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from .communication_op import *
from .parallel_state import *
from .utils import *
@@ -1,51 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/communication_op.py
from typing import Any, Dict, Optional, Union
import torch
import torch.distributed
from .parallel_state import get_tp_group, get_sp_group
def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
"""All-reduce the input tensor across model parallel group."""
return get_tp_group().all_reduce(input_)
def tensor_model_parallel_all_gather(input_: torch.Tensor,
dim: int = -1) -> torch.Tensor:
"""All-gather the input tensor across model parallel group."""
return get_tp_group().all_gather(input_, dim)
def tensor_model_parallel_gather(input_: torch.Tensor,
dst: int = 0,
dim: int = -1) -> Optional[torch.Tensor]:
"""Gather the input tensor across model parallel group."""
return get_tp_group().gather(input_, dst, dim)
def broadcast_tensor_dict(tensor_dict: Optional[Dict[Any, Union[torch.Tensor,
Any]]] = None,
src: int = 0):
if not torch.distributed.is_initialized():
return tensor_dict
return get_tp_group().broadcast_tensor_dict(tensor_dict, src)
# TODO: remove model, make it sequence_parallel
def sequence_model_parallel_all_to_all_4D(input_: torch.Tensor,
scatter_dim: int = 2,
gather_dim: int = 1) -> torch.Tensor:
"""All-to-all communication of 4D tensors (e.g. QKV matrices) across sequence parallel group."""
return get_sp_group().all_to_all_4D(input_, scatter_dim, gather_dim)
def sequence_model_parallel_all_gather(input_: torch.Tensor,
dim: int = -1) -> torch.Tensor:
"""All-gather the input tensor across model parallel group."""
return get_sp_group().all_gather(input_, dim)
@@ -1,184 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/base_device_communicator.py
from typing import Optional
import torch
import torch.distributed as dist
from torch.distributed import ProcessGroup
from einops import rearrange
class DeviceCommunicatorBase:
"""
Base class for device-specific communicator.
It can use the `cpu_group` to initialize the communicator.
If the device has PyTorch integration (PyTorch can recognize its
communication backend), the `device_group` will also be given.
"""
def __init__(self,
cpu_group: ProcessGroup,
device: Optional[torch.device] = None,
device_group: Optional[ProcessGroup] = None,
unique_name: str = ""):
self.device = device or torch.device("cpu")
self.cpu_group = cpu_group
self.device_group = device_group
self.unique_name = unique_name
self.rank = dist.get_rank(cpu_group)
self.world_size = dist.get_world_size(cpu_group)
self.ranks = dist.get_process_group_ranks(cpu_group)
self.global_rank = dist.get_rank()
self.global_world_size = dist.get_world_size()
self.rank_in_group = dist.get_group_rank(self.cpu_group,
self.global_rank)
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
dist.all_reduce(input_, group=self.device_group)
return input_
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
if dim < 0:
# Convert negative dim to positive.
dim += input_.dim()
input_size = input_.size()
# NOTE: we have to use concat-style all-gather here,
# stack-style all-gather has compatibility issues with
# torch.compile . see https://github.com/pytorch/pytorch/issues/138795
output_size = (input_size[0] * self.world_size, ) + input_size[1:]
# Allocate output tensor.
output_tensor = torch.empty(output_size,
dtype=input_.dtype,
device=input_.device)
# All-gather.
dist.all_gather_into_tensor(output_tensor,
input_,
group=self.device_group)
# Reshape
output_tensor = output_tensor.reshape((self.world_size, ) + input_size)
output_tensor = output_tensor.movedim(0, dim)
output_tensor = output_tensor.reshape(input_size[:dim] +
(self.world_size *
input_size[dim], ) +
input_size[dim + 1:])
return output_tensor
def gather(self,
input_: torch.Tensor,
dst: int = 0,
dim: int = -1) -> Optional[torch.Tensor]:
"""
NOTE: We assume that the input tensor is on the same device across
all the ranks.
NOTE: `dst` is the local rank of the destination rank.
"""
world_size = self.world_size
assert -input_.dim() <= dim < input_.dim(), (
f"Invalid dim ({dim}) for input tensor with shape {input_.size()}")
if dim < 0:
# Convert negative dim to positive.
dim += input_.dim()
# Allocate output tensor.
if self.rank_in_group == dst:
gather_list = [torch.empty_like(input_) for _ in range(world_size)]
else:
gather_list = None
# Gather.
torch.distributed.gather(input_,
gather_list,
dst=self.ranks[dst],
group=self.device_group)
if self.rank_in_group == dst:
output_tensor = torch.cat(gather_list, dim=dim)
else:
output_tensor = None
return output_tensor
def all_to_all_4D(self,
input_: torch.Tensor,
scatter_dim: int = 2,
gather_dim: int = 1) -> torch.Tensor:
"""Specialized all-to-all operation for 4D tensors (e.g., for QKV matrices).
Args:
input_ (torch.Tensor): 4D input tensor to be scattered and gathered.
scatter_dim (int, optional): Dimension along which to scatter. Defaults to 2.
gather_dim (int, optional): Dimension along which to gather. Defaults to 1.
Returns:
torch.Tensor: Output tensor after all-to-all operation.
"""
# Bypass the function if we are using only 1 GPU.
if self.world_size == 1:
return input_
assert input_.dim() == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
if scatter_dim == 2 and gather_dim == 1:
# input: (bs, seqlen/P, hc, hs) output: (bs, seqlen, hc/P, hs)
bs, shard_seqlen, hc, hs = input_.shape
seqlen = shard_seqlen * self.world_size
shard_hc = hc // self.world_size
# Reshape and transpose for scattering
input_t = (input_.reshape(bs, shard_seqlen, self.world_size, shard_hc, hs).transpose(0, 2).contiguous())
output = torch.empty_like(input_t)
torch.distributed.all_to_all_single(output, input_t, group=self.device_group)
torch.cuda.synchronize()
# Reshape and transpose back
output = output.reshape(seqlen, bs, shard_hc, hs).transpose(0, 1).contiguous().reshape(bs, seqlen, shard_hc, hs)
return output
elif scatter_dim == 1 and gather_dim == 2:
# input: (bs, seqlen, hc/P, hs) output: (bs, seqlen/P, hc, hs)
bs, seqlen, shard_hc, hs = input_.shape
hc = shard_hc * self.world_size
shard_seqlen = seqlen // self.world_size
# Reshape and transpose for scattering
input_t = (input_.reshape(bs, self.world_size, shard_seqlen, shard_hc,
hs).transpose(0,
3).transpose(0,
1).contiguous().reshape(self.world_size, shard_hc,
shard_seqlen, bs, hs))
output = torch.empty_like(input_t)
torch.distributed.all_to_all_single(output, input_t, group=self.device_group)
torch.cuda.synchronize()
# Reshape and transpose back
output = output.reshape(hc, shard_seqlen, bs, hs).transpose(0, 2).contiguous().reshape(bs, shard_seqlen, hc, hs)
return output
else:
raise RuntimeError("scatter_dim must be 1 or 2 and gather_dim must be 1 or 2")
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank."""
if dst is None:
dst = (self.rank_in_group + 1) % self.world_size
torch.distributed.send(tensor, self.ranks[dst], self.device_group)
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: Optional[int] = None) -> torch.Tensor:
"""Receives a tensor from the source rank."""
"""NOTE: `src` is the local rank of the source rank."""
if src is None:
src = (self.rank_in_group - 1) % self.world_size
tensor = torch.empty(size, dtype=dtype, device=self.device)
torch.distributed.recv(tensor, self.ranks[src], self.device_group)
return tensor
def destroy(self):
pass
@@ -1,109 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/cuda_communicator.py
from typing import Optional
import torch
from torch.distributed import ProcessGroup
from .base_device_communicator import DeviceCommunicatorBase
class CudaCommunicator(DeviceCommunicatorBase):
def __init__(self,
cpu_group: ProcessGroup,
device: Optional[torch.device] = None,
device_group: Optional[ProcessGroup] = None,
unique_name: str = ""):
super().__init__(cpu_group, device, device_group, unique_name)
if "pp" in unique_name:
# pipeline parallel does not need custom allreduce
use_custom_allreduce = False
else:
# from vllm.distributed.parallel_state import (
# _ENABLE_CUSTOM_ALL_REDUCE)
# TODO(will): bring in the custom allreduce from vLLM
use_custom_allreduce = False
use_pynccl = True
self.use_pynccl = use_pynccl
self.use_custom_allreduce = use_custom_allreduce
# lazy import to avoid documentation build error
# from vllm.distributed.device_communicators.custom_all_reduce import (
# CustomAllreduce)
from fastvideo.v1.distributed.device_communicators.pynccl import (
PyNcclCommunicator)
self.pynccl_comm: Optional[PyNcclCommunicator] = None
if use_pynccl and self.world_size > 1:
self.pynccl_comm = PyNcclCommunicator(
group=self.cpu_group,
device=self.device,
)
# TODO(will): bring in the custom allreduce from vLLM
self.ca_comm: Optional[CustomAllreduce] = None
if use_custom_allreduce and self.world_size > 1:
# Initialize a custom fast all-reduce implementation.
self.ca_comm = CustomAllreduce(
group=self.cpu_group,
device=self.device,
)
def all_reduce(self, input_):
# always try custom allreduce first,
# and then pynccl.
ca_comm = self.ca_comm
if ca_comm is not None and not ca_comm.disabled and \
ca_comm.should_custom_ar(input_):
out = ca_comm.custom_all_reduce(input_)
assert out is not None
return out
pynccl_comm = self.pynccl_comm
assert pynccl_comm is not None
out = pynccl_comm.all_reduce(input_)
if out is None:
# fall back to the default all-reduce using PyTorch.
# this usually happens during testing.
# when we run the model, allreduce only happens for the TP
# group, where we always have either custom allreduce or pynccl.
out = input_.clone()
torch.distributed.all_reduce(out, group=self.device_group)
return out
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank."""
if dst is None:
dst = (self.rank_in_group + 1) % self.world_size
pynccl_comm = self.pynccl_comm
if pynccl_comm is not None and not pynccl_comm.disabled:
pynccl_comm.send(tensor, dst)
else:
torch.distributed.send(tensor, self.ranks[dst], self.device_group)
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: Optional[int] = None) -> torch.Tensor:
"""Receives a tensor from the source rank."""
"""NOTE: `src` is the local rank of the source rank."""
if src is None:
src = (self.rank_in_group - 1) % self.world_size
tensor = torch.empty(size, dtype=dtype, device=self.device)
pynccl_comm = self.pynccl_comm
if pynccl_comm is not None and not pynccl_comm.disabled:
pynccl_comm.recv(tensor, src)
else:
torch.distributed.recv(tensor, self.ranks[src], self.device_group)
return tensor
def destroy(self):
if self.pynccl_comm is not None:
self.pynccl_comm = None
if self.ca_comm is not None:
self.ca_comm = None
@@ -1,218 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/pynccl.py
from typing import Optional, Union
# ===================== import region =====================
import torch
import torch.distributed as dist
from torch.distributed import ProcessGroup, ReduceOp
from fastvideo.v1.distributed.device_communicators.pynccl_wrapper import (
NCCLLibrary, buffer_type, cudaStream_t, ncclComm_t, ncclDataTypeEnum,
ncclRedOpTypeEnum, ncclUniqueId)
from fastvideo.v1.distributed.utils import StatelessProcessGroup
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import current_stream
logger = init_logger(__name__)
class PyNcclCommunicator:
def __init__(
self,
group: Union[ProcessGroup, StatelessProcessGroup],
device: Union[int, str, torch.device],
library_path: Optional[str] = None,
):
"""
Args:
group: the process group to work on. If None, it will use the
default process group.
device: the device to bind the PyNcclCommunicator to. If None,
it will be bind to f"cuda:{local_rank}".
library_path: the path to the NCCL library. If None, it will
use the default library path.
It is the caller's responsibility to make sure each communicator
is bind to a unique device.
"""
if not isinstance(group, StatelessProcessGroup):
assert dist.is_initialized()
assert dist.get_backend(group) != dist.Backend.NCCL, (
"PyNcclCommunicator should be attached to a non-NCCL group.")
# note: this rank is the rank in the group
self.rank = dist.get_rank(group)
self.world_size = dist.get_world_size(group)
else:
self.rank = group.rank
self.world_size = group.world_size
self.group = group
# if world_size == 1, no need to create communicator
if self.world_size == 1:
self.available = False
self.disabled = True
return
try:
self.nccl = NCCLLibrary(library_path)
except Exception:
# disable because of missing NCCL library
# e.g. in a non-GPU environment
self.available = False
self.disabled = True
return
self.available = True
self.disabled = False
logger.info("FastVideo is using nccl==%s", self.nccl.ncclGetVersion())
if self.rank == 0:
# get the unique id from NCCL
self.unique_id = self.nccl.ncclGetUniqueId()
else:
# construct an empty unique id
self.unique_id = ncclUniqueId()
if not isinstance(group, StatelessProcessGroup):
tensor = torch.ByteTensor(list(self.unique_id.internal))
ranks = dist.get_process_group_ranks(group)
# arg `src` in `broadcast` is the global rank
dist.broadcast(tensor, src=ranks[0], group=group)
byte_list = tensor.tolist()
for i, byte in enumerate(byte_list):
self.unique_id.internal[i] = byte
else:
self.unique_id = group.broadcast_obj(self.unique_id, src=0)
if isinstance(device, int):
device = torch.device(f"cuda:{device}")
elif isinstance(device, str):
device = torch.device(device)
# now `device` is a `torch.device` object
assert isinstance(device, torch.device)
self.device = device
# nccl communicator and stream will use this device
# `torch.cuda.device` is a context manager that changes the
# current cuda device to the specified one
with torch.cuda.device(device):
self.comm: ncclComm_t = self.nccl.ncclCommInitRank(
self.world_size, self.unique_id, self.rank)
stream = current_stream()
# A small all_reduce for warmup.
data = torch.zeros(1, device=device)
self.all_reduce(data)
stream.synchronize()
del data
def all_reduce(self,
in_tensor: torch.Tensor,
op: ReduceOp = ReduceOp.SUM,
stream=None) -> torch.Tensor:
if self.disabled:
return None
# nccl communicator created on a specific device
# will only work on tensors on the same device
# otherwise it will cause "illegal memory access"
assert in_tensor.device == self.device, (
f"this nccl communicator is created to work on {self.device}, "
f"but the input tensor is on {in_tensor.device}")
out_tensor = torch.empty_like(in_tensor)
if stream is None:
stream = current_stream()
self.nccl.ncclAllReduce(buffer_type(in_tensor.data_ptr()),
buffer_type(out_tensor.data_ptr()),
in_tensor.numel(),
ncclDataTypeEnum.from_torch(in_tensor.dtype),
ncclRedOpTypeEnum.from_torch(op), self.comm,
cudaStream_t(stream.cuda_stream))
return out_tensor
def all_gather(self,
output_tensor: torch.Tensor,
input_tensor: torch.Tensor,
stream=None):
if self.disabled:
return
# nccl communicator created on a specific device
# will only work on tensors on the same device
# otherwise it will cause "illegal memory access"
assert input_tensor.device == self.device, (
f"this nccl communicator is created to work on {self.device}, "
f"but the input tensor is on {input_tensor.device}")
if stream is None:
stream = current_stream()
self.nccl.ncclAllGather(
buffer_type(input_tensor.data_ptr()),
buffer_type(output_tensor.data_ptr()), input_tensor.numel(),
ncclDataTypeEnum.from_torch(input_tensor.dtype), self.comm,
cudaStream_t(stream.cuda_stream))
def reduce_scatter(self,
output_tensor: torch.Tensor,
input_tensor: torch.Tensor,
op: ReduceOp = ReduceOp.SUM,
stream=None):
if self.disabled:
return
# nccl communicator created on a specific device
# will only work on tensors on the same device
# otherwise it will cause "illegal memory access"
assert input_tensor.device == self.device, (
f"this nccl communicator is created to work on {self.device}, "
f"but the input tensor is on {input_tensor.device}")
if stream is None:
stream = current_stream()
self.nccl.ncclReduceScatter(
buffer_type(input_tensor.data_ptr()),
buffer_type(output_tensor.data_ptr()), output_tensor.numel(),
ncclDataTypeEnum.from_torch(input_tensor.dtype),
ncclRedOpTypeEnum.from_torch(op), self.comm,
cudaStream_t(stream.cuda_stream))
def send(self, tensor: torch.Tensor, dst: int, stream=None):
if self.disabled:
return
assert tensor.device == self.device, (
f"this nccl communicator is created to work on {self.device}, "
f"but the input tensor is on {tensor.device}")
if stream is None:
stream = current_stream()
self.nccl.ncclSend(buffer_type(tensor.data_ptr()), tensor.numel(),
ncclDataTypeEnum.from_torch(tensor.dtype), dst,
self.comm, cudaStream_t(stream.cuda_stream))
def recv(self, tensor: torch.Tensor, src: int, stream=None):
if self.disabled:
return
assert tensor.device == self.device, (
f"this nccl communicator is created to work on {self.device}, "
f"but the input tensor is on {tensor.device}")
if stream is None:
stream = current_stream()
self.nccl.ncclRecv(buffer_type(tensor.data_ptr()), tensor.numel(),
ncclDataTypeEnum.from_torch(tensor.dtype), src,
self.comm, cudaStream_t(stream.cuda_stream))
def broadcast(self, tensor: torch.Tensor, src: int, stream=None):
if self.disabled:
return
assert tensor.device == self.device, (
f"this nccl communicator is created to work on {self.device}, "
f"but the input tensor is on {tensor.device}")
if stream is None:
stream = current_stream()
if src == self.rank:
sendbuff = buffer_type(tensor.data_ptr())
# NCCL requires the sender also to have a receive buffer
recvbuff = buffer_type(tensor.data_ptr())
else:
sendbuff = buffer_type()
recvbuff = buffer_type(tensor.data_ptr())
self.nccl.ncclBroadcast(sendbuff, recvbuff, tensor.numel(),
ncclDataTypeEnum.from_torch(tensor.dtype), src,
self.comm, cudaStream_t(stream.cuda_stream))
@@ -1,342 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/pynccl_wrapper.py
# This file is a pure Python wrapper for the NCCL library.
# The main purpose is to use NCCL combined with CUDA graph.
# Before writing this script, we tried the following approach:
# 1. We tried to use `cupy`, it calls NCCL correctly, but `cupy` itself
# often gets stuck when initializing the NCCL communicator.
# 2. We tried to use `torch.distributed`, but `torch.distributed.all_reduce`
# contains many other potential cuda APIs, that are not allowed during
# capturing the CUDA graph. For further details, please check
# https://discuss.pytorch.org/t/pytorch-cudagraph-with-nccl-operation-failed/ .
#
# Another rejected idea is to write a C/C++ binding for NCCL. It is usually
# doable, but we often encounter issues related with nccl versions, and need
# to switch between different versions of NCCL. See
# https://github.com/NVIDIA/nccl/issues/1234 for more details.
# A C/C++ binding is not flexible enough to handle this. It requires
# recompilation of the code every time we want to switch between different
# versions. This current implementation, with a **pure** Python wrapper, is
# more flexible. We can easily switch between different versions of NCCL by
# changing the environment variable `VLLM_NCCL_SO_PATH`, or the `so_file`
# variable in the code.
import ctypes
import platform
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
import torch
from torch.distributed import ReduceOp
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import find_nccl_library
logger = init_logger(__name__)
# === export types and functions from nccl to Python ===
# for the original nccl definition, please check
# https://github.com/NVIDIA/nccl/blob/master/src/nccl.h.in
ncclResult_t = ctypes.c_int
ncclComm_t = ctypes.c_void_p
class ncclUniqueId(ctypes.Structure):
_fields_ = [("internal", ctypes.c_byte * 128)]
cudaStream_t = ctypes.c_void_p
buffer_type = ctypes.c_void_p
ncclDataType_t = ctypes.c_int
class ncclDataTypeEnum:
ncclInt8 = 0
ncclChar = 0
ncclUint8 = 1
ncclInt32 = 2
ncclInt = 2
ncclUint32 = 3
ncclInt64 = 4
ncclUint64 = 5
ncclFloat16 = 6
ncclHalf = 6
ncclFloat32 = 7
ncclFloat = 7
ncclFloat64 = 8
ncclDouble = 8
ncclBfloat16 = 9
ncclNumTypes = 10
@classmethod
def from_torch(cls, dtype: torch.dtype) -> int:
if dtype == torch.int8:
return cls.ncclInt8
if dtype == torch.uint8:
return cls.ncclUint8
if dtype == torch.int32:
return cls.ncclInt32
if dtype == torch.int64:
return cls.ncclInt64
if dtype == torch.float16:
return cls.ncclFloat16
if dtype == torch.float32:
return cls.ncclFloat32
if dtype == torch.float64:
return cls.ncclFloat64
if dtype == torch.bfloat16:
return cls.ncclBfloat16
raise ValueError(f"Unsupported dtype: {dtype}")
ncclRedOp_t = ctypes.c_int
class ncclRedOpTypeEnum:
ncclSum = 0
ncclProd = 1
ncclMax = 2
ncclMin = 3
ncclAvg = 4
ncclNumOps = 5
@classmethod
def from_torch(cls, op: ReduceOp) -> int:
if op == ReduceOp.SUM:
return cls.ncclSum
if op == ReduceOp.PRODUCT:
return cls.ncclProd
if op == ReduceOp.MAX:
return cls.ncclMax
if op == ReduceOp.MIN:
return cls.ncclMin
if op == ReduceOp.AVG:
return cls.ncclAvg
raise ValueError(f"Unsupported op: {op}")
@dataclass
class Function:
name: str
restype: Any
argtypes: List[Any]
class NCCLLibrary:
exported_functions = [
# const char* ncclGetErrorString(ncclResult_t result)
Function("ncclGetErrorString", ctypes.c_char_p, [ncclResult_t]),
# ncclResult_t ncclGetVersion(int *version);
Function("ncclGetVersion", ncclResult_t,
[ctypes.POINTER(ctypes.c_int)]),
# ncclResult_t ncclGetUniqueId(ncclUniqueId* uniqueId);
Function("ncclGetUniqueId", ncclResult_t,
[ctypes.POINTER(ncclUniqueId)]),
# ncclResult_t ncclCommInitRank(
# ncclComm_t* comm, int nranks, ncclUniqueId commId, int rank);
# note that ncclComm_t is a pointer type, so the first argument
# is a pointer to a pointer
Function("ncclCommInitRank", ncclResult_t, [
ctypes.POINTER(ncclComm_t), ctypes.c_int, ncclUniqueId,
ctypes.c_int
]),
# ncclResult_t ncclAllReduce(
# const void* sendbuff, void* recvbuff, size_t count,
# ncclDataType_t datatype, ncclRedOp_t op, ncclComm_t comm,
# cudaStream_t stream);
# note that cudaStream_t is a pointer type, so the last argument
# is a pointer
Function("ncclAllReduce", ncclResult_t, [
buffer_type, buffer_type, ctypes.c_size_t, ncclDataType_t,
ncclRedOp_t, ncclComm_t, cudaStream_t
]),
# ncclResult_t ncclAllGather(
# const void* sendbuff, void* recvbuff, size_t count,
# ncclDataType_t datatype, ncclComm_t comm,
# cudaStream_t stream);
# note that cudaStream_t is a pointer type, so the last argument
# is a pointer
Function("ncclAllGather", ncclResult_t, [
buffer_type, buffer_type, ctypes.c_size_t, ncclDataType_t,
ncclComm_t, cudaStream_t
]),
# ncclResult_t ncclReduceScatter(
# const void* sendbuff, void* recvbuff, size_t count,
# ncclDataType_t datatype, ncclRedOp_t op, ncclComm_t comm,
# cudaStream_t stream);
# note that cudaStream_t is a pointer type, so the last argument
# is a pointer
Function("ncclReduceScatter", ncclResult_t, [
buffer_type, buffer_type, ctypes.c_size_t, ncclDataType_t,
ncclRedOp_t, ncclComm_t, cudaStream_t
]),
# ncclResult_t ncclSend(
# const void* sendbuff, size_t count, ncclDataType_t datatype,
# int dest, ncclComm_t comm, cudaStream_t stream);
Function("ncclSend", ncclResult_t, [
buffer_type, ctypes.c_size_t, ncclDataType_t, ctypes.c_int,
ncclComm_t, cudaStream_t
]),
# ncclResult_t ncclRecv(
# void* recvbuff, size_t count, ncclDataType_t datatype,
# int src, ncclComm_t comm, cudaStream_t stream);
Function("ncclRecv", ncclResult_t, [
buffer_type, ctypes.c_size_t, ncclDataType_t, ctypes.c_int,
ncclComm_t, cudaStream_t
]),
# ncclResult_t ncclBroadcast(
# const void* sendbuff, void* recvbuff, size_t count,
# ncclDataType_t datatype, int root, ncclComm_t comm,
# cudaStream_t stream);
Function("ncclBroadcast", ncclResult_t, [
buffer_type, buffer_type, ctypes.c_size_t, ncclDataType_t,
ctypes.c_int, ncclComm_t, cudaStream_t
]),
# be cautious! this is a collective call, it will block until all
# processes in the communicator have called this function.
# because Python object destruction can happen in random order,
# it is better not to call it at all.
# ncclResult_t ncclCommDestroy(ncclComm_t comm);
Function("ncclCommDestroy", ncclResult_t, [ncclComm_t]),
]
# class attribute to store the mapping from the path to the library
# to avoid loading the same library multiple times
path_to_library_cache: Dict[str, Any] = {}
# class attribute to store the mapping from library path
# to the corresponding dictionary
path_to_dict_mapping: Dict[str, Dict[str, Any]] = {}
def __init__(self, so_file: Optional[str] = None):
so_file = so_file or find_nccl_library()
try:
if so_file not in NCCLLibrary.path_to_dict_mapping:
lib = ctypes.CDLL(so_file)
NCCLLibrary.path_to_library_cache[so_file] = lib
self.lib = NCCLLibrary.path_to_library_cache[so_file]
except Exception as e:
logger.error(
"Failed to load NCCL library from %s ."
"It is expected if you are not running on NVIDIA/AMD GPUs."
"Otherwise, the nccl library might not exist, be corrupted "
"or it does not support the current platform %s."
"If you already have the library, please set the "
"environment variable VLLM_NCCL_SO_PATH"
" to point to the correct nccl library path.", so_file,
platform.platform())
raise e
if so_file not in NCCLLibrary.path_to_dict_mapping:
_funcs: Dict[str, Any] = {}
for func in NCCLLibrary.exported_functions:
f = getattr(self.lib, func.name)
f.restype = func.restype
f.argtypes = func.argtypes
_funcs[func.name] = f
NCCLLibrary.path_to_dict_mapping[so_file] = _funcs
self._funcs = NCCLLibrary.path_to_dict_mapping[so_file]
def ncclGetErrorString(self, result: ncclResult_t) -> str:
return self._funcs["ncclGetErrorString"](result).decode("utf-8")
def NCCL_CHECK(self, result: ncclResult_t) -> None:
if result != 0:
error_str = self.ncclGetErrorString(result)
raise RuntimeError(f"NCCL error: {error_str}")
def ncclGetVersion(self) -> str:
version = ctypes.c_int()
self.NCCL_CHECK(self._funcs["ncclGetVersion"](ctypes.byref(version)))
version_str = str(version.value)
# something like 21903 --> "2.19.3"
major = version_str[0].lstrip("0")
minor = version_str[1:3].lstrip("0")
patch = version_str[3:].lstrip("0")
return f"{major}.{minor}.{patch}"
def ncclGetUniqueId(self) -> ncclUniqueId:
unique_id = ncclUniqueId()
self.NCCL_CHECK(self._funcs["ncclGetUniqueId"](
ctypes.byref(unique_id)))
return unique_id
def ncclCommInitRank(self, world_size: int, unique_id: ncclUniqueId,
rank: int) -> ncclComm_t:
comm = ncclComm_t()
self.NCCL_CHECK(self._funcs["ncclCommInitRank"](ctypes.byref(comm),
world_size, unique_id,
rank))
return comm
def ncclAllReduce(self, sendbuff: buffer_type, recvbuff: buffer_type,
count: int, datatype: int, op: int, comm: ncclComm_t,
stream: cudaStream_t) -> None:
# `datatype` actually should be `ncclDataType_t`
# and `op` should be `ncclRedOp_t`
# both are aliases of `ctypes.c_int`
# when we pass int to a function, it will be converted to `ctypes.c_int`
# by ctypes automatically
self.NCCL_CHECK(self._funcs["ncclAllReduce"](sendbuff, recvbuff, count,
datatype, op, comm,
stream))
def ncclReduceScatter(self, sendbuff: buffer_type, recvbuff: buffer_type,
count: int, datatype: int, op: int, comm: ncclComm_t,
stream: cudaStream_t) -> None:
# `datatype` actually should be `ncclDataType_t`
# and `op` should be `ncclRedOp_t`
# both are aliases of `ctypes.c_int`
# when we pass int to a function, it will be converted to `ctypes.c_int`
# by ctypes automatically
self.NCCL_CHECK(self._funcs["ncclReduceScatter"](sendbuff, recvbuff,
count, datatype, op,
comm, stream))
def ncclAllGather(self, sendbuff: buffer_type, recvbuff: buffer_type,
count: int, datatype: int, comm: ncclComm_t,
stream: cudaStream_t) -> None:
# `datatype` actually should be `ncclDataType_t`
# which is an aliases of `ctypes.c_int`
# when we pass int to a function, it will be converted to `ctypes.c_int`
# by ctypes automatically
self.NCCL_CHECK(self._funcs["ncclAllGather"](sendbuff, recvbuff, count,
datatype, comm, stream))
def ncclSend(self, sendbuff: buffer_type, count: int, datatype: int,
dest: int, comm: ncclComm_t, stream: cudaStream_t) -> None:
self.NCCL_CHECK(self._funcs["ncclSend"](sendbuff, count, datatype,
dest, comm, stream))
def ncclRecv(self, recvbuff: buffer_type, count: int, datatype: int,
src: int, comm: ncclComm_t, stream: cudaStream_t) -> None:
self.NCCL_CHECK(self._funcs["ncclRecv"](recvbuff, count, datatype, src,
comm, stream))
def ncclBroadcast(self, sendbuff: buffer_type, recvbuff: buffer_type,
count: int, datatype: int, root: int, comm: ncclComm_t,
stream: cudaStream_t) -> None:
self.NCCL_CHECK(self._funcs["ncclBroadcast"](sendbuff, recvbuff, count,
datatype, root, comm,
stream))
def ncclCommDestroy(self, comm: ncclComm_t) -> None:
self.NCCL_CHECK(self._funcs["ncclCommDestroy"](comm))
__all__ = [
"NCCLLibrary", "ncclDataTypeEnum", "ncclRedOpTypeEnum", "ncclUniqueId",
"ncclComm_t", "cudaStream_t", "buffer_type"
]
File diff suppressed because it is too large Load Diff
-204
View File
@@ -1,204 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/utils.py
# Copyright 2023 The vLLM team.
# Adapted from
# https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/core/tensor_parallel/utils.py
# Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved.
import dataclasses
import pickle
import time
from collections import deque
from typing import Any, Deque, Dict, Optional, Sequence, Tuple
import torch
from torch.distributed import ProcessGroup, TCPStore
from torch.distributed.distributed_c10d import (Backend, PrefixStore,
_get_default_timeout,
is_nccl_available)
from torch.distributed.rendezvous import rendezvous
import fastvideo.v1.envs as envs
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def ensure_divisibility(numerator, denominator):
"""Ensure that numerator is divisible by the denominator."""
assert numerator % denominator == 0, "{} is not divisible by {}".format(
numerator, denominator)
def divide(numerator, denominator):
"""Ensure that numerator is divisible by the denominator and return
the division value."""
ensure_divisibility(numerator, denominator)
return numerator // denominator
def split_tensor_along_last_dim(
tensor: torch.Tensor,
num_partitions: int,
contiguous_split_chunks: bool = False,
) -> Sequence[torch.Tensor]:
""" Split a tensor along its last dimension.
Arguments:
tensor: input tensor.
num_partitions: number of partitions to split the tensor
contiguous_split_chunks: If True, make each chunk contiguous
in memory.
Returns:
A list of Tensors
"""
# Get the size and dimension.
last_dim = tensor.dim() - 1
last_dim_size = divide(tensor.size()[last_dim], num_partitions)
# Split.
tensor_list = torch.split(tensor, last_dim_size, dim=last_dim)
# NOTE: torch.split does not create contiguous tensors by default.
if contiguous_split_chunks:
return tuple(chunk.contiguous() for chunk in tensor_list)
return tensor_list
@dataclasses.dataclass
class StatelessProcessGroup:
"""A dataclass to hold a metadata store, and the rank, world_size of the
group. Only use it to communicate metadata between processes.
For data-plane communication, create NCCL-related objects.
"""
rank: int
world_size: int
store: torch._C._distributed_c10d.Store
data_expiration_seconds: int = 3600 # 1 hour
# dst rank -> counter
send_dst_counter: Dict[int, int] = dataclasses.field(default_factory=dict)
# src rank -> counter
recv_src_counter: Dict[int, int] = dataclasses.field(default_factory=dict)
broadcast_send_counter: int = 0
broadcast_recv_src_counter: Dict[int, int] = dataclasses.field(
default_factory=dict)
# A deque to store the data entries, with key and timestamp.
entries: Deque[Tuple[str,
float]] = dataclasses.field(default_factory=deque)
def __post_init__(self):
assert self.rank < self.world_size
self.send_dst_counter = {i: 0 for i in range(self.world_size)}
self.recv_src_counter = {i: 0 for i in range(self.world_size)}
self.broadcast_recv_src_counter = {
i: 0
for i in range(self.world_size)
}
def send_obj(self, obj: Any, dst: int):
"""Send an object to a destination rank."""
self.expire_data()
key = f"send_to/{dst}/{self.send_dst_counter[dst]}"
self.store.set(key, pickle.dumps(obj))
self.send_dst_counter[dst] += 1
self.entries.append((key, time.time()))
def expire_data(self):
"""Expire data that is older than `data_expiration_seconds` seconds."""
while self.entries:
# check the oldest entry
key, timestamp = self.entries[0]
if time.time() - timestamp > self.data_expiration_seconds:
self.store.delete_key(key)
self.entries.popleft()
else:
break
def recv_obj(self, src: int) -> Any:
"""Receive an object from a source rank."""
obj = pickle.loads(
self.store.get(
f"send_to/{self.rank}/{self.recv_src_counter[src]}"))
self.recv_src_counter[src] += 1
return obj
def broadcast_obj(self, obj: Optional[Any], src: int) -> Any:
"""Broadcast an object from a source rank to all other ranks.
It does not clean up after all ranks have received the object.
Use it for limited times, e.g., for initialization.
"""
if self.rank == src:
self.expire_data()
key = (f"broadcast_from/{src}/"
f"{self.broadcast_send_counter}")
self.store.set(key, pickle.dumps(obj))
self.broadcast_send_counter += 1
self.entries.append((key, time.time()))
return obj
else:
key = (f"broadcast_from/{src}/"
f"{self.broadcast_recv_src_counter[src]}")
recv_obj = pickle.loads(self.store.get(key))
self.broadcast_recv_src_counter[src] += 1
return recv_obj
def all_gather_obj(self, obj: Any) -> list[Any]:
"""All gather an object from all ranks."""
gathered_objs = []
for i in range(self.world_size):
if i == self.rank:
gathered_objs.append(obj)
self.broadcast_obj(obj, src=self.rank)
else:
recv_obj = self.broadcast_obj(None, src=i)
gathered_objs.append(recv_obj)
return gathered_objs
def barrier(self):
"""A barrier to synchronize all ranks."""
for i in range(self.world_size):
if i == self.rank:
self.broadcast_obj(None, src=self.rank)
else:
self.broadcast_obj(None, src=i)
@staticmethod
def create(
host: str,
port: int,
rank: int,
world_size: int,
data_expiration_seconds: int = 3600,
) -> "StatelessProcessGroup":
"""A replacement for `torch.distributed.init_process_group` that does not
pollute the global state.
If we have process A and process B called `torch.distributed.init_process_group`
to form a group, and then we want to form another group with process A, B, C,
D, it is not possible in PyTorch, because process A and process B have already
formed a group, and process C and process D cannot join that group. This
function is a workaround for this issue.
`torch.distributed.init_process_group` is a global call, while this function
is a stateless call. It will return a `StatelessProcessGroup` object that can be
used for exchanging metadata. With this function, process A and process B
can call `StatelessProcessGroup.create` to form a group, and then process A, B,
C, and D can call `StatelessProcessGroup.create` to form another group.
""" # noqa
store = TCPStore(
host_name=host,
port=port,
world_size=world_size,
is_master=(rank == 0),
)
return StatelessProcessGroup(
rank=rank,
world_size=world_size,
store=store,
data_expiration_seconds=data_expiration_seconds)
-225
View File
@@ -1,225 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm
# https://github.com/vllm-project/vllm/blob/b382a7f28f739f3b120e5495fd029089d0399428/vllm/envs.py
# Copyright 2023 The vLLM Authors.
# Copyright 2023 The FastVideo Authors.
import os
import tempfile
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional
if TYPE_CHECKING:
FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60
FASTVIDEO_NCCL_SO_PATH: Optional[str] = None
LD_LIBRARY_PATH: Optional[str] = None
FASTVIDEO_USE_TRITON_FLASH_ATTN: bool = False
FASTVIDEO_FLASH_ATTN_VERSION: Optional[int] = None
LOCAL_RANK: int = 0
CUDA_VISIBLE_DEVICES: Optional[str] = None
FASTVIDEO_CACHE_ROOT: str = os.path.expanduser("~/.cache/fastvideo")
FASTVIDEO_CONFIG_ROOT: str = os.path.expanduser("~/.config/fastvideo")
FASTVIDEO_CONFIGURE_LOGGING: int = 1
FASTVIDEO_LOGGING_LEVEL: str = "INFO"
FASTVIDEO_LOGGING_PREFIX: str = ""
FASTVIDEO_LOGGING_CONFIG_PATH: Optional[str] = None
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: Optional[str] = None
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: Optional[str] = None
NVCC_THREADS: Optional[str] = None
CMAKE_BUILD_TYPE: Optional[str] = None
VERBOSE: bool = False
FASTVIDEO_SERVER_DEV_MODE: bool = False
def get_default_cache_root():
return os.getenv(
"XDG_CACHE_HOME",
os.path.join(os.path.expanduser("~"), ".cache"),
)
def get_default_config_root():
return os.getenv(
"XDG_CONFIG_HOME",
os.path.join(os.path.expanduser("~"), ".config"),
)
def maybe_convert_int(value: Optional[str]) -> Optional[int]:
if value is None:
return None
return int(value)
# The begin-* and end* here are used by the documentation generator
# to extract the used env vars.
# begin-env-vars-definition
environment_variables: Dict[str, Callable[[], Any]] = {
# ================== Installation Time Env Vars ==================
# Target device of FastVideo, supporting [cuda (by default),
# rocm, neuron, cpu, openvino]
"FASTVIDEO_TARGET_DEVICE":
lambda: os.getenv("FASTVIDEO_TARGET_DEVICE", "cuda"),
# Maximum number of compilation jobs to run in parallel.
# By default this is the number of CPUs
"MAX_JOBS":
lambda: os.getenv("MAX_JOBS", None),
# Number of threads to use for nvcc
# By default this is 1.
# If set, `MAX_JOBS` will be reduced to avoid oversubscribing the CPU.
"NVCC_THREADS":
lambda: os.getenv("NVCC_THREADS", None),
# If set, fastvideo will use precompiled binaries (*.so)
"FASTVIDEO_USE_PRECOMPILED":
lambda: bool(os.environ.get("FASTVIDEO_USE_PRECOMPILED")) or bool(
os.environ.get("FASTVIDEO_PRECOMPILED_WHEEL_LOCATION")),
# CMake build type
# If not set, defaults to "Debug" or "RelWithDebInfo"
# Available options: "Debug", "Release", "RelWithDebInfo"
"CMAKE_BUILD_TYPE":
lambda: os.getenv("CMAKE_BUILD_TYPE"),
# If set, fastvideo will print verbose logs during installation
"VERBOSE":
lambda: bool(int(os.getenv('VERBOSE', '0'))),
# Root directory for FASTVIDEO configuration files
# Defaults to `~/.config/fastvideo` unless `XDG_CONFIG_HOME` is set
# Note that this not only affects how fastvideo finds its configuration files
# during runtime, but also affects how fastvideo installs its configuration
# files during **installation**.
"FASTVIDEO_CONFIG_ROOT":
lambda: os.path.expanduser(
os.getenv(
"FASTVIDEO_CONFIG_ROOT",
os.path.join(get_default_config_root(), "fastvideo"),
)),
# ================== Runtime Env Vars ==================
# Root directory for FASTVIDEO cache files
# Defaults to `~/.cache/fastvideo` unless `XDG_CACHE_HOME` is set
"FASTVIDEO_CACHE_ROOT":
lambda: os.path.expanduser(
os.getenv(
"FASTVIDEO_CACHE_ROOT",
os.path.join(get_default_cache_root(), "fastvideo"),
)),
# Interval in seconds to log a warning message when the ring buffer is full
"FASTVIDEO_RINGBUFFER_WARNING_INTERVAL":
lambda: int(os.environ.get("FASTVIDEO_RINGBUFFER_WARNING_INTERVAL", "60")),
# Path to the NCCL library file. It is needed because nccl>=2.19 brought
# by PyTorch contains a bug: https://github.com/NVIDIA/nccl/issues/1234
"FASTVIDEO_NCCL_SO_PATH":
lambda: os.environ.get("FASTVIDEO_NCCL_SO_PATH", None),
# when `FASTVIDEO_NCCL_SO_PATH` is not set, fastvideo will try to find the nccl
# library file in the locations specified by `LD_LIBRARY_PATH`
"LD_LIBRARY_PATH":
lambda: os.environ.get("LD_LIBRARY_PATH", None),
# flag to control if fastvideo should use triton flash attention
"FASTVIDEO_USE_TRITON_FLASH_ATTN":
lambda: (os.environ.get("FASTVIDEO_USE_TRITON_FLASH_ATTN", "True").lower() in
("true", "1")),
# Force fastvideo to use a specific flash-attention version (2 or 3), only valid
# when using the flash-attention backend.
"FASTVIDEO_FLASH_ATTN_VERSION":
lambda: maybe_convert_int(os.environ.get("FASTVIDEO_FLASH_ATTN_VERSION", None)),
# Internal flag to enable Dynamo fullgraph capture
"FASTVIDEO_TEST_DYNAMO_FULLGRAPH_CAPTURE":
lambda: bool(
os.environ.get("FASTVIDEO_TEST_DYNAMO_FULLGRAPH_CAPTURE", "1") != "0"),
# local rank of the process in the distributed setting, used to determine
# the GPU device id
"LOCAL_RANK":
lambda: int(os.environ.get("LOCAL_RANK", "0")),
# used to control the visible devices in the distributed setting
"CUDA_VISIBLE_DEVICES":
lambda: os.environ.get("CUDA_VISIBLE_DEVICES", None),
# timeout for each iteration in the engine
"FASTVIDEO_ENGINE_ITERATION_TIMEOUT_S":
lambda: int(os.environ.get("FASTVIDEO_ENGINE_ITERATION_TIMEOUT_S", "60")),
# Logging configuration
# If set to 0, fastvideo will not configure logging
# If set to 1, fastvideo will configure logging using the default configuration
# or the configuration file specified by FASTVIDEO_LOGGING_CONFIG_PATH
"FASTVIDEO_CONFIGURE_LOGGING":
lambda: int(os.getenv("FASTVIDEO_CONFIGURE_LOGGING", "1")),
"FASTVIDEO_LOGGING_CONFIG_PATH":
lambda: os.getenv("FASTVIDEO_LOGGING_CONFIG_PATH"),
# this is used for configuring the default logging level
"FASTVIDEO_LOGGING_LEVEL":
lambda: os.getenv("FASTVIDEO_LOGGING_LEVEL", "INFO"),
# if set, FASTVIDEO_LOGGING_PREFIX will be prepended to all log messages
"FASTVIDEO_LOGGING_PREFIX":
lambda: os.getenv("FASTVIDEO_LOGGING_PREFIX", ""),
# Trace function calls
# If set to 1, fastvideo will trace function calls
# Useful for debugging
"FASTVIDEO_TRACE_FUNCTION":
lambda: int(os.getenv("FASTVIDEO_TRACE_FUNCTION", "0")),
# Backend for attention computation
# Available options:
# - "TORCH_SDPA": use torch.nn.MultiheadAttention
# - "FLASH_ATTN": use FlashAttention
# - "XFORMERS": use XFormers
# - "ROCM_FLASH": use ROCmFlashAttention
# - "FLASHINFER": use flashinfer
"FASTVIDEO_ATTENTION_BACKEND":
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
# Use dedicated multiprocess context for workers.
# Both spawn and fork work
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "fork"),
# Enables torch profiler if set. Path to the directory where torch profiler
# traces are saved. Note that it must be an absolute path.
"FASTVIDEO_TORCH_PROFILER_DIR":
lambda: (None if os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", None) is None else os
.path.expanduser(os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", "."))),
# If set, fastvideo will run in development mode, which will enable
# some additional endpoints for developing and debugging,
# e.g. `/reset_prefix_cache`
"FASTVIDEO_SERVER_DEV_MODE":
lambda: bool(int(os.getenv("FASTVIDEO_SERVER_DEV_MODE", "0"))),
}
# end-env-vars-definition
def __getattr__(name: str):
# lazy evaluation of environment variables
if name in environment_variables:
return environment_variables[name]()
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
def __dir__():
return list(environment_variables.keys())
-471
View File
@@ -1,471 +0,0 @@
# Copyright 2023-2024 SGLang Team
# Adapted from SGLang server_args.py
# Copyright 2024-2025 FastVideo Team
# 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.
# ==============================================================================
"""The arguments of FastVideo Inference."""
import argparse
import dataclasses
from fastvideo.v1.utils import FlexibleArgumentParser
from typing import List, Optional
@dataclasses.dataclass
class InferenceArgs:
# Model and path configuration
model_path: str
# HuggingFace specific parameters
trust_remote_code: bool = False
revision: Optional[str] = None
# Parallelism
tp_size: int = 1
sp_size: int = 1
dist_timeout: Optional[int] = None # timeout for torch.distributed
# Video generation parameters
height: int = 720
width: int = 1280
num_frames: int = 117
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: int = 7
output_type: str = "pil"
# Model configuration
precision: str = "bf16"
# VAE configurationi
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = False
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_len: int = 256
hidden_state_skip_layer: int = 2
# Secondary text encoder
text_encoder_precision_2: str = "fp16"
text_len_2: int = 77
# Flow Matching parameters
flow_solver: str = "euler"
denoise_type: str = "flow"
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
# Scheduler options
scheduler_type: str = "euler"
neg_prompt: Optional[str] = None
num_videos: int = 1
fps: int = 24
use_cpu_offload: bool = False
disable_autocast: bool = False
# Logging
log_level: str = "info"
# Kernel backend
attention_backend: Optional[str] = None
# Inference parameters
prompt: Optional[str] = None
prompt_path: Optional[str] = None
output_path: str = "outputs/"
seed: int = 1024
device_str: Optional[str] = None
device = None
def __post_init__(self):
pass
@staticmethod
def add_cli_args(parser: argparse.ArgumentParser):
parser.add_argument(
"--use-v1-text-encoder",
action="store_true",
help="Use the v1 text encoder",
)
parser.add_argument(
"--use-v1-vae",
action="store_true",
help="Use the v1 vae",
)
parser.add_argument(
"--use-v1-transformer",
action="store_true",
help="Use the v1 transformer",
)
# Model and path configuration
parser.add_argument(
"--model-path",
type=str,
help="The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
required=True,
)
parser.add_argument(
"--dit-weight",
type=str,
help="Path to the DiT model weights",
)
parser.add_argument(
"--model-dir",
type=str,
help="Directory containing StepVideo model",
)
# HuggingFace specific parameters
parser.add_argument(
"--trust-remote-code",
action="store_true",
default=InferenceArgs.trust_remote_code,
help="Trust remote code when loading HuggingFace models",
)
parser.add_argument(
"--revision",
type=str,
default=InferenceArgs.revision,
help="The specific model version to use (can be a branch name, tag name, or commit id)",
)
# Parallelism
parser.add_argument(
"--tensor-parallel-size",
"--tp-size",
type=int,
default=InferenceArgs.tp_size,
help="The tensor parallelism size.",
)
parser.add_argument(
"--sequence-parallel-size",
"--sp-size",
type=int,
default=InferenceArgs.sp_size,
help="The sequence parallelism size.",
)
parser.add_argument(
"--dist-timeout",
type=int,
default=InferenceArgs.dist_timeout,
help="Set timeout for torch.distributed initialization.",
)
# Video generation parameters
parser.add_argument(
"--height",
type=int,
default=InferenceArgs.height,
help="Height of generated video",
)
parser.add_argument(
"--width",
type=int,
default=InferenceArgs.width,
help="Width of generated video",
)
parser.add_argument(
"--num-frames",
type=int,
default=InferenceArgs.num_frames,
help="Number of frames to generate",
)
parser.add_argument(
"--num-inference-steps",
type=int,
default=InferenceArgs.num_inference_steps,
help="Number of inference steps",
)
parser.add_argument(
"--guidance-scale",
type=float,
default=InferenceArgs.guidance_scale,
help="Guidance scale for classifier-free guidance",
)
parser.add_argument(
"--guidance-rescale",
type=float,
default=InferenceArgs.guidance_rescale,
help="Guidance rescale for classifier-free guidance",
)
parser.add_argument(
"--embedded-cfg-scale",
type=float,
default=InferenceArgs.embedded_cfg_scale,
help="Embedded CFG scale",
)
parser.add_argument(
"--flow-shift",
"--shift",
type=int,
default=InferenceArgs.flow_shift,
help="Flow shift parameter",
)
parser.add_argument(
"--output-type",
type=str,
default=InferenceArgs.output_type,
choices=["pil"],
help="Output type for the generated video",
)
parser.add_argument(
"--precision",
type=str,
default=InferenceArgs.precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for the model",
)
# VAE configuration
parser.add_argument(
"--vae-precision",
type=str,
default=InferenceArgs.vae_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for VAE",
)
parser.add_argument(
"--vae-tiling",
action="store_true",
default=InferenceArgs.vae_tiling,
help="Enable VAE tiling",
)
parser.add_argument(
"--vae-sp",
action="store_true",
help="Enable VAE spatial parallelism",
)
parser.add_argument(
"--text-encoder-precision",
type=str,
default=InferenceArgs.text_encoder_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for text encoder",
)
parser.add_argument(
"--text-len",
type=int,
default=InferenceArgs.text_len,
help="Maximum text length",
)
# Secondary text encoder
parser.add_argument(
"--text-encoder-precision-2",
type=str,
default=InferenceArgs.text_encoder_precision_2,
choices=["fp32", "fp16", "bf16"],
help="Precision for secondary text encoder",
)
parser.add_argument(
"--text-len-2",
type=int,
default=InferenceArgs.text_len_2,
help="Maximum secondary text length",
)
# Flow Matching parameters
parser.add_argument(
"--flow-solver",
type=str,
default=InferenceArgs.flow_solver,
help="Solver for flow matching",
)
parser.add_argument(
"--denoise-type",
type=str,
default=InferenceArgs.denoise_type,
help="Denoise type for noised inputs",
)
# STA (Spatial-Temporal Attention) parameters
parser.add_argument(
"--mask-strategy-file-path",
type=str,
help="Path to mask strategy JSON file for STA",
)
parser.add_argument(
"--enable-torch-compile",
action="store_true",
help="Use torch.compile for speeding up STA inference without teacache",
)
# Scheduler options
parser.add_argument(
"--scheduler-type",
type=str,
default=InferenceArgs.scheduler_type,
help="Type of scheduler to use",
)
# HunYuan specific parameters
parser.add_argument(
"--neg-prompt",
type=str,
default=InferenceArgs.neg_prompt,
help="Negative prompt for sampling",
)
parser.add_argument(
"--num-videos",
type=int,
default=InferenceArgs.num_videos,
help="Number of videos to generate per prompt",
)
parser.add_argument(
"--fps",
type=int,
default=InferenceArgs.fps,
help="Frames per second for output video",
)
parser.add_argument(
"--use-cpu-offload",
action="store_true",
help="Use CPU offload for the model load",
)
parser.add_argument(
"--disable-autocast",
action="store_true",
help="Disable autocast for denoising loop and vae decoding in pipeline sampling",
)
# Logging
parser.add_argument(
"--log-level",
type=str,
default=InferenceArgs.log_level,
help="The logging level of all loggers.",
)
# Kernel backend
parser.add_argument(
"--attention-backend",
type=str,
choices=["flashinfer", "triton", "torch_native"],
default=InferenceArgs.attention_backend,
help="Choose the kernels for attention layers.",
)
# Inference parameters
parser.add_argument(
"--prompt",
type=str,
help="Text prompt for video generation",
)
parser.add_argument(
"--prompt-path",
type=str,
help="Path to a text file containing the prompt",
)
parser.add_argument(
"--output-path",
type=str,
default=InferenceArgs.output_path,
help="Directory to save generated videos",
)
parser.add_argument(
"--seed",
type=int,
default=InferenceArgs.seed,
help="Random seed for reproducibility",
)
@classmethod
def from_cli_args(cls, args: argparse.Namespace):
args.tp_size = args.tensor_parallel_size
args.sp_size = args.sequence_parallel_size
args.flow_shift = getattr(args, "shift", args.flow_shift)
# Get all fields from the dataclass
attrs = [attr.name for attr in dataclasses.fields(cls)]
# Create a dictionary of attribute values, with defaults for missing attributes
kwargs = {}
for attr in attrs:
# Convert snake_case attribute name to kebab-case CLI argument name
cli_attr = attr.replace('_', '-')
# Handle renamed attributes or those with multiple CLI names
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
# Use getattr with default value from the dataclass for potentially missing attributes
else:
default_value = getattr(cls, attr, None)
kwargs[attr] = getattr(args, attr, default_value)
return cls(**kwargs)
def check_inference_args(self):
"""Validate inference arguments for consistency"""
# Validate VAE spatial parallelism with VAE tiling
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.")
assert self.prompt is not None or self.prompt_path is not None, "Either prompt or prompt_path must be provided"
assert self.prompt_path.endswith(".txt"), "prompt_path must be a text file"
_inference_args = None
def prepare_inference_args(argv: List[str]) -> InferenceArgs:
"""
Prepare the inference arguments from the command line arguments.
Args:
argv: The command line arguments. Typically, it should be `sys.argv[1:]`
to ensure compatibility with `parse_args` when no arguments are passed.
Returns:
The inference arguments.
"""
parser = FlexibleArgumentParser()
InferenceArgs.add_cli_args(parser)
raw_args = parser.parse_args(argv)
inference_args = InferenceArgs.from_cli_args(raw_args)
inference_args.check_inference_args()
global _inference_args
_inference_args = inference_args
return inference_args
def get_inference_args() -> InferenceArgs:
global _inference_args
return _inference_args
class DeprecatedAction(argparse.Action):
def __init__(self, option_strings, dest, nargs=0, **kwargs):
super(DeprecatedAction, self).__init__(
option_strings, dest, nargs=nargs, **kwargs
)
def __call__(self, parser, namespace, values, option_string=None):
raise ValueError(self.help)
-213
View File
@@ -1,213 +0,0 @@
"""
Inference module for diffusion models.
This module provides classes and functions for running inference with diffusion models.
"""
import os
import time
import torch
from typing import Any, Dict
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.pipelines import ComposedPipelineBase, build_pipeline
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.logger import init_logger
# TODO(will): remove, check if this is hunyuan specific
from fastvideo.v1.utils import align_to
# TODO(will): remove, move this to hunyuan stage
from fastvideo.v1.pipelines.implementations.hunyuan.constants import NEGATIVE_PROMPT
logger = init_logger(__name__)
class InferenceEngine:
"""
Engine for running inference with diffusion models.
"""
def __init__(
self,
pipeline: ComposedPipelineBase,
inference_args: InferenceArgs,
):
"""
Initialize the inference engine.
Args:
pipeline: The pipeline to use for inference.
inference_args: The inference arguments.
default_negative_prompt: The default negative prompt to use.
"""
self.pipeline = pipeline
self.inference_args = inference_args
# TODO(will): this is a hack to get the default negative prompt
self.default_negative_prompt = NEGATIVE_PROMPT
@classmethod
def create_engine(
cls,
inference_args: InferenceArgs,
) -> "InferenceEngine":
"""
Create an inference engine with the specified arguments.
Args:
inference_args: The inference arguments.
model_loader_cls: The model loader class to use. If None, it will be
determined from the model type.
pipeline_type: The type of pipeline to create. If None, it will be
determined from the model type.
Returns:
The created inference engine.
Raises:
ValueError: If the model type is not recognized or if the pipeline type
is not recognized.
"""
logger.info(f"Building pipeline...")
# TODO(will): I don't really like this api.
# it should be something closer to pipeline_cls.from_pretrained(...)
# this way for training we can just do pipeline_cls.from_pretrained(
# checkpoint_path) and have it handle everything.
# TODO(Peiyuan): Then maybe we should only pass in model path and device, not the entire inference args?
pipeline = build_pipeline(inference_args)
logger.info(f"Pipeline Ready")
# Create the inference engine
return cls(pipeline, inference_args)
def run(
self,
prompt: str,
inference_args: InferenceArgs,
) -> Dict[str, Any]:
"""
Run inference with the pipeline.
Args:
prompt: The prompt to use for generation.
negative_prompt: The negative prompt to use. If None, the default will be used.
seed: The random seed to use. If None, a random seed will be used.
**kwargs: Additional arguments to pass to the pipeline.
Returns:
A dictionary containing the generated videos and metadata.
"""
out_dict = dict()
num_videos_per_prompt = inference_args.num_videos
seed = inference_args.seed
height = inference_args.height
width = inference_args.width
video_length = inference_args.num_frames
negative_prompt = inference_args.neg_prompt
infer_steps = inference_args.num_inference_steps
guidance_scale = inference_args.guidance_scale
flow_shift = inference_args.flow_shift
embedded_guidance_scale = inference_args.embedded_cfg_scale
# ========================================================================
# Arguments: target_width, target_height, target_video_length
# ========================================================================
if width <= 0 or height <= 0 or video_length <= 0:
raise ValueError(
f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={video_length}"
)
if (video_length - 1) % 4 != 0:
raise ValueError(f"`video_length-1` must be a multiple of 4, got {video_length}")
logger.info(f"Input (height, width, video_length) = ({height}, {width}, {video_length})")
target_height = align_to(height, 16)
target_width = align_to(width, 16)
target_video_length = video_length
out_dict["size"] = (target_height, target_width, target_video_length)
# ========================================================================
# Arguments: prompt, new_prompt, negative_prompt
# ========================================================================
if not isinstance(prompt, str):
raise TypeError(f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
# negative prompt
if negative_prompt is None or negative_prompt == "":
negative_prompt = self.default_negative_prompt
if not isinstance(negative_prompt, str):
raise TypeError(f"`negative_prompt` must be a string, but got {type(negative_prompt)}")
negative_prompt = negative_prompt.strip()
# TODO(PY): move to hunyuan stage
latents_size = [(video_length - 1) // 4 + 1, height // 8, width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# ========================================================================
# Print infer args
# ========================================================================
debug_str = f"""
height: {target_height}
width: {target_width}
video_length: {target_video_length}
prompt: {prompt}
neg_prompt: {negative_prompt}
seed: {seed}
infer_steps: {infer_steps}
num_videos_per_prompt: {num_videos_per_prompt}
guidance_scale: {guidance_scale}
n_tokens: {n_tokens}
flow_shift: {flow_shift}
embedded_guidance_scale: {embedded_guidance_scale}"""
logger.info(debug_str)
# return
# sp_group = get_sp_group()
# local_rank = sp_group.rank
device = torch.device(inference_args.device_str)
batch = ForwardBatch(
prompt=prompt,
negative_prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
height=inference_args.height,
width=inference_args.width,
num_frames=inference_args.num_frames,
num_inference_steps=inference_args.num_inference_steps,
guidance_scale=inference_args.guidance_scale,
# generator=generator,
eta=0.0,
n_tokens=n_tokens,
data_type="video" if inference_args.num_frames > 1 else "image",
device=device,
extra={}, # Any additional parameters
)
print('===============================================')
print(batch)
print('===============================================')
print('===============================================')
print(inference_args)
# ========================================================================
# Pipeline inference
# ========================================================================
start_time = time.time()
samples = self.pipeline.forward(
batch=batch,
inference_args=inference_args,
)[0]
# TODO(will): fix and move to hunyuan stage
# out_dict["seeds"] = batch.seeds
out_dict["samples"] = samples
out_dict["prompts"] = prompt
gen_time = time.time() - start_time
logger.info(f"Success, time: {gen_time}")
return out_dict
View File
-351
View File
@@ -1,351 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Custom activation functions."""
import math
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from vllm.distributed import (divide, get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size)
from vllm.model_executor.custom_op import CustomOp
from vllm.model_executor.utils import set_weight_attrs
from fastvideo.v1.platforms import current_platform
@CustomOp.register("fatrelu_and_mul")
class FatreluAndMul(CustomOp):
"""An activation function for FATReLU.
The function computes x -> FATReLU(x[:d]) * x[d:] where
d = x.shape[-1] // 2.
This is used in openbmb/MiniCPM-S-1B-sft.
Shapes:
x: (num_tokens, 2 * d) or (batch_size, seq_len, 2 * d)
return: (num_tokens, d) or (batch_size, seq_len, d)
"""
def __init__(self, threshold: float = 0.):
super().__init__()
self.threshold = threshold
if current_platform.is_cuda_alike():
self.op = torch.ops._C.fatrelu_and_mul
elif current_platform.is_cpu():
self._forward_method = self.forward_native
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
x1 = x[..., :d]
x2 = x[..., d:]
x1 = F.threshold(x1, self.threshold, 0.0)
return x1 * x2
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
output_shape = (x.shape[:-1] + (d, ))
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
self.op(out, x, self.threshold)
return out
@CustomOp.register("silu_and_mul")
class SiluAndMul(CustomOp):
"""An activation function for SwiGLU.
The function computes x -> silu(x[:d]) * x[d:] where d = x.shape[-1] // 2.
Shapes:
x: (num_tokens, 2 * d) or (batch_size, seq_len, 2 * d)
return: (num_tokens, d) or (batch_size, seq_len, d)
"""
def __init__(self):
super().__init__()
if current_platform.is_cuda_alike() or current_platform.is_cpu():
self.op = torch.ops._C.silu_and_mul
elif current_platform.is_xpu():
from vllm._ipex_ops import ipex_ops
self.op = ipex_ops.silu_and_mul
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
d = x.shape[-1] // 2
return F.silu(x[..., :d]) * x[..., d:]
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
output_shape = (x.shape[:-1] + (d, ))
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
self.op(out, x)
return out
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
output_shape = (x.shape[:-1] + (d, ))
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
self.op(out, x)
return out
@CustomOp.register("mul_and_silu")
class MulAndSilu(CustomOp):
"""An activation function for SwiGLU.
The function computes x -> x[:d] * silu(x[d:]) where d = x.shape[-1] // 2.
Shapes:
x: (num_tokens, 2 * d) or (batch_size, seq_len, 2 * d)
return: (num_tokens, d) or (batch_size, seq_len, d)
"""
def __init__(self):
super().__init__()
if current_platform.is_cuda_alike():
self.op = torch.ops._C.mul_and_silu
elif current_platform.is_xpu():
from vllm._ipex_ops import ipex_ops
self.op = ipex_ops.silu_and_mul
elif current_platform.is_cpu():
self._forward_method = self.forward_native
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
d = x.shape[-1] // 2
return x[..., :d] * F.silu(x[..., d:])
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
output_shape = (x.shape[:-1] + (d, ))
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
self.op(out, x)
return out
# TODO implement forward_xpu for MulAndSilu
# def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
@CustomOp.register("gelu_and_mul")
class GeluAndMul(CustomOp):
"""An activation function for GeGLU.
The function computes x -> GELU(x[:d]) * x[d:] where d = x.shape[-1] // 2.
Shapes:
x: (batch_size, seq_len, 2 * d) or (num_tokens, 2 * d)
return: (batch_size, seq_len, d) or (num_tokens, d)
"""
def __init__(self, approximate: str = "none"):
super().__init__()
self.approximate = approximate
if approximate not in ("none", "tanh"):
raise ValueError(f"Unknown approximate mode: {approximate}")
if current_platform.is_cuda_alike() or current_platform.is_cpu():
if approximate == "none":
self.op = torch.ops._C.gelu_and_mul
elif approximate == "tanh":
self.op = torch.ops._C.gelu_tanh_and_mul
elif current_platform.is_xpu():
from vllm._ipex_ops import ipex_ops
if approximate == "none":
self.op = ipex_ops.gelu_and_mul
else:
self.op = ipex_ops.gelu_tanh_and_mul
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
d = x.shape[-1] // 2
return F.gelu(x[..., :d], approximate=self.approximate) * x[..., d:]
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
output_shape = (x.shape[:-1] + (d, ))
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
self.op(out, x)
return out
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
output_shape = (x.shape[:-1] + (d, ))
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
self.op(out, x)
return out
def extra_repr(self) -> str:
return f'approximate={repr(self.approximate)}'
@CustomOp.register("gelu_new")
class NewGELU(CustomOp):
def __init__(self):
super().__init__()
if current_platform.is_cuda_alike() or current_platform.is_cpu():
self.op = torch.ops._C.gelu_new
elif current_platform.is_xpu():
from vllm._ipex_ops import ipex_ops
self.op = ipex_ops.gelu_new
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
c = math.sqrt(2.0 / math.pi)
return 0.5 * x * (1.0 + torch.tanh(c *
(x + 0.044715 * torch.pow(x, 3.0))))
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
self.op(out, x)
return out
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
return self.op(x)
@CustomOp.register("gelu_fast")
class FastGELU(CustomOp):
def __init__(self):
super().__init__()
if current_platform.is_cuda_alike() or current_platform.is_cpu():
self.op = torch.ops._C.gelu_fast
elif current_platform.is_xpu():
from vllm._ipex_ops import ipex_ops
self.op = ipex_ops.gelu_fast
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
return 0.5 * x * (1.0 + torch.tanh(x * 0.7978845608 *
(1.0 + 0.044715 * x * x)))
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
self.op(out, x)
return out
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
return self.op(x)
@CustomOp.register("quick_gelu")
class QuickGELU(CustomOp):
# https://github.com/huggingface/transformers/blob/main/src/transformers/activations.py#L90
def __init__(self):
super().__init__()
if current_platform.is_cuda_alike() or current_platform.is_cpu():
self.op = torch.ops._C.gelu_quick
elif current_platform.is_xpu():
from vllm._ipex_ops import ipex_ops
self.op = ipex_ops.gelu_quick
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
return x * torch.sigmoid(1.702 * x)
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
self.op(out, x)
return out
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
self.op(out, x)
return out
# TODO implement forward_xpu for QuickGELU
# def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
@CustomOp.register("relu2")
class ReLUSquaredActivation(CustomOp):
"""
Applies the relu^2 activation introduced in https://arxiv.org/abs/2109.08668v2
"""
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
return torch.square(F.relu(x))
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
return self.forward_native(x)
class ScaledActivation(nn.Module):
"""An activation function with post-scale parameters.
This is used for some quantization methods like AWQ.
"""
def __init__(
self,
act_module: nn.Module,
intermediate_size: int,
input_is_parallel: bool = True,
params_dtype: Optional[torch.dtype] = None,
):
super().__init__()
self.act = act_module
self.input_is_parallel = input_is_parallel
if input_is_parallel:
tp_size = get_tensor_model_parallel_world_size()
intermediate_size_per_partition = divide(intermediate_size,
tp_size)
else:
intermediate_size_per_partition = intermediate_size
if params_dtype is None:
params_dtype = torch.get_default_dtype()
self.scales = nn.Parameter(
torch.empty(intermediate_size_per_partition, dtype=params_dtype))
set_weight_attrs(self.scales, {"weight_loader": self.weight_loader})
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.act(x) / self.scales
def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
param_data = param.data
if self.input_is_parallel:
tp_rank = get_tensor_model_parallel_rank()
shard_size = param_data.shape[0]
start_idx = tp_rank * shard_size
loaded_weight = loaded_weight.narrow(0, start_idx, shard_size)
assert param_data.shape == loaded_weight.shape
param_data.copy_(loaded_weight)
_ACTIVATION_REGISTRY = {
"gelu": nn.GELU,
"gelu_fast": FastGELU,
"gelu_new": NewGELU,
"gelu_pytorch_tanh": lambda: nn.GELU(approximate="tanh"),
"relu": nn.ReLU,
"relu2": ReLUSquaredActivation,
"silu": nn.SiLU,
"quick_gelu": QuickGELU,
}
def get_act_fn(act_fn_name: str) -> nn.Module:
"""Get an activation function by name."""
act_fn_name = act_fn_name.lower()
if act_fn_name not in _ACTIVATION_REGISTRY:
raise ValueError(
f"Activation function {act_fn_name!r} is not supported.")
return _ACTIVATION_REGISTRY[act_fn_name]()
_ACTIVATION_AND_MUL_REGISTRY = {
"gelu": GeluAndMul,
"silu": SiluAndMul,
}
def get_act_and_mul_fn(act_fn_name: str) -> nn.Module:
"""Get an activation-and-mul (i.e. SiluAndMul) function by name."""
act_fn_name = act_fn_name.lower()
if act_fn_name not in _ACTIVATION_AND_MUL_REGISTRY:
raise ValueError(
f"Activation function {act_fn_name!r} is not supported.")
return _ACTIVATION_AND_MUL_REGISTRY[act_fn_name]()
View File
-308
View File
@@ -1,308 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Custom normalization layers."""
from typing import Optional, Tuple, Union
import torch
import torch.nn as nn
from vllm.model_executor.custom_op import CustomOp
@CustomOp.register("rms_norm")
class RMSNorm(CustomOp):
"""Root mean square normalization.
Computes x -> w * x / sqrt(E[x^2] + eps) where w is the learned weight.
Refer to https://arxiv.org/abs/1910.07467
"""
def __init__(
self,
hidden_size: int,
eps: float = 1e-6,
var_hidden_size: Optional[int] = None,
has_weight: bool = True,
) -> None:
super().__init__()
self.hidden_size = hidden_size
self.variance_epsilon = eps
self.variance_size_override = (None if var_hidden_size == hidden_size
else var_hidden_size)
self.has_weight = has_weight
self.weight = torch.ones(hidden_size)
if self.has_weight:
self.weight = nn.Parameter(self.weight)
def forward_native(
self,
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""PyTorch-native implementation equivalent to forward()."""
orig_dtype = x.dtype
x = x.to(torch.float32)
if residual is not None:
x = x + residual.to(torch.float32)
residual = x.to(orig_dtype)
hidden_size = x.shape[-1]
if hidden_size != self.hidden_size:
raise ValueError("Expected hidden_size to be "
f"{self.hidden_size}, but found: {hidden_size}")
if self.variance_size_override is None:
x_var = x
else:
if hidden_size < self.variance_size_override:
raise ValueError(
"Expected hidden_size to be at least "
f"{self.variance_size_override}, but found: {hidden_size}")
x_var = x[:, :, :self.variance_size_override]
variance = x_var.pow(2).mean(dim=-1, keepdim=True)
x = x * torch.rsqrt(variance + self.variance_epsilon)
x = x.to(orig_dtype)
if self.has_weight:
x = x * self.weight
if residual is None:
return x
else:
return x, residual
def forward_cuda(
self,
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
if self.variance_size_override is not None:
return self.forward_native(x, residual)
from vllm import _custom_ops as ops
if residual is not None:
ops.fused_add_rms_norm(
x,
residual,
self.weight.data,
self.variance_epsilon,
)
return x, residual
out = torch.empty_like(x)
ops.rms_norm(
out,
x,
self.weight.data,
self.variance_epsilon,
)
return out
def forward_hpu(
self,
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
from vllm_hpu_extension.ops import HPUFusedRMSNorm
if HPUFusedRMSNorm is None:
return self.forward_native(x, residual)
if residual is not None:
orig_shape = x.shape
residual += x.view(residual.shape)
# Note: HPUFusedRMSNorm requires 3D tensors as inputs
x = HPUFusedRMSNorm.apply(residual, self.weight,
self.variance_epsilon)
return x.view(orig_shape), residual
x = HPUFusedRMSNorm.apply(x, self.weight, self.variance_epsilon)
return x
def forward_xpu(
self,
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
if self.variance_size_override is not None:
return self.forward_native(x, residual)
from vllm._ipex_ops import ipex_ops as ops
if residual is not None:
ops.fused_add_rms_norm(
x,
residual,
self.weight.data,
self.variance_epsilon,
)
return x, residual
return ops.rms_norm(
x,
self.weight.data,
self.variance_epsilon,
)
def extra_repr(self) -> str:
s = f"hidden_size={self.weight.data.size(0)}"
s += f", eps={self.variance_epsilon}"
return s
@CustomOp.register("gemma_rms_norm")
class GemmaRMSNorm(CustomOp):
"""RMS normalization for Gemma.
Two differences from the above RMSNorm:
1. x * (1 + w) instead of x * w.
2. (x * w).to(orig_dtype) instead of x.to(orig_dtype) * w.
"""
def __init__(
self,
hidden_size: int,
eps: float = 1e-6,
) -> None:
super().__init__()
self.weight = nn.Parameter(torch.zeros(hidden_size))
self.variance_epsilon = eps
@staticmethod
def forward_static(
weight: torch.Tensor,
variance_epsilon: float,
x: torch.Tensor,
residual: Optional[torch.Tensor],
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""PyTorch-native implementation equivalent to forward()."""
orig_dtype = x.dtype
if residual is not None:
x = x + residual
residual = x
x = x.float()
variance = x.pow(2).mean(dim=-1, keepdim=True)
x = x * torch.rsqrt(variance + variance_epsilon)
# Llama does x.to(float16) * w whilst Gemma is (x * w).to(float16)
# See https://github.com/huggingface/transformers/pull/29402
x = x * (1.0 + weight.float())
x = x.to(orig_dtype)
return x if residual is None else (x, residual)
def forward_native(
self,
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""PyTorch-native implementation equivalent to forward()."""
return self.forward_static(self.weight.data, self.variance_epsilon, x,
residual)
def forward_cuda(
self,
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
if torch.compiler.is_compiling():
return self.forward_native(x, residual)
if not getattr(self, "_is_compiled", False):
self.forward_static = torch.compile( # type: ignore
self.forward_static)
self._is_compiled = True
return self.forward_native(x, residual)
class ScaleResidual(nn.Module):
"""
Applies gated residual connection.
"""
def __init__(self):
super().__init__()
def forward(self, residual: torch.Tensor, x: torch.Tensor, gate: torch.Tensor) -> torch.Tensor:
"""Apply gated residual connection."""
return residual + x * gate
class ScaleResidualLayerNormScaleShift(nn.Module):
"""
Fused operation that combines:
1. Gated residual connection
2. LayerNorm
3. Scale and shift operations
This reduces memory bandwidth by combining memory-bound operations.
"""
def __init__(
self,
hidden_size: int,
norm_type: str = "rms",
eps: float = 1e-6,
elementwise_affine: bool = False,
dtype: torch.dtype = torch.float32,
):
super().__init__()
if norm_type == "rms":
self.norm = RMSNorm(hidden_size, has_weight=elementwise_affine, eps=eps, dtype=dtype)
elif norm_type == "layer":
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=elementwise_affine, eps=eps, dtype=dtype)
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
def forward(
self,
residual: torch.Tensor,
x: torch.Tensor,
gate: torch.Tensor,
shift: torch.Tensor,
scale: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Apply gated residual connection, followed by layernorm and scale/shift in a single fused operation.
Returns:
Tuple containing:
- normalized and modulated output
- residual value (value after residual connection but before normalization)
"""
# Apply residual connection with gating
residual_output = residual + x * gate
# Apply normalization
normalized = self.norm(residual_output)
# Apply scale and shift
modulated = normalized * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
return modulated, residual_output
class LayerNormScaleShift(nn.Module):
"""
Fused operation that combines LayerNorm with scale and shift operations.
This reduces memory bandwidth by combining memory-bound operations.
"""
def __init__(
self,
hidden_size: int,
norm_type: str = "rms",
eps: float = 1e-6,
elementwise_affine: bool = False,
dtype: torch.dtype = torch.float32,
):
super().__init__()
if norm_type == "rms":
self.norm = RMSNorm(hidden_size, has_weight=elementwise_affine, eps=eps)
elif norm_type == "layer":
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=elementwise_affine, eps=eps, dtype=dtype)
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
def forward(self, x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
"""Apply layernorm followed by scale and shift in a single fused operation."""
normalized = self.norm(x)
return normalized * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
File diff suppressed because it is too large Load Diff
-46
View File
@@ -1,46 +0,0 @@
from fastvideo.v1.layers.linear import ReplicatedLinear
from fastvideo.v1.layers.activation import get_act_fn
import torch
import torch.nn as nn
from typing import Optional
class MLP(nn.Module):
"""
MLP for DiT blocks, NO gated linear units
TODO: add Tensor Parallel
"""
def __init__(
self,
input_dim: int,
mlp_hidden_dim: int,
output_dim: int = None,
bias: bool = True,
act_type: str = "gelu_pytorch_tanh",
dtype: Optional[torch.dtype] = None,
):
super().__init__()
self.fc_in = ReplicatedLinear(
input_dim,
mlp_hidden_dim, # For activation functions like SiLU that need 2x width
bias=bias,
params_dtype=dtype
)
self.act = get_act_fn(act_type)
if output_dim is None:
output_dim = input_dim
self.fc_out = ReplicatedLinear(
mlp_hidden_dim,
output_dim,
bias=bias,
params_dtype=dtype
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x, _ = self.fc_in(x)
x = self.act(x)
x, _ = self.fc_out(x)
return x
-541
View File
@@ -1,541 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from
# https://github.com/huggingface/transformers/blob/v4.33.2/src/transformers/models/llama/modeling_llama.py
# Copyright 2023 The vLLM team.
# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.
#
# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
# and OPT implementations in this library. It has been modified from its
# original forms to accommodate minor architectural differences compared
# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
#
# 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.
"""Rotary Positional Embeddings."""
import math
from typing import Any, Dict, List, Optional, Tuple, Union
import torch
import torch.nn as nn
from transformers import PretrainedConfig
from vllm.model_executor.custom_op import CustomOp
from fastvideo.v1.distributed.parallel_state import get_sp_group
def _rotate_neox(x: torch.Tensor) -> torch.Tensor:
x1 = x[..., :x.shape[-1] // 2]
x2 = x[..., x.shape[-1] // 2:]
return torch.cat((-x2, x1), dim=-1)
def _rotate_gptj(x: torch.Tensor) -> torch.Tensor:
x1 = x[..., ::2]
x2 = x[..., 1::2]
x = torch.stack((-x2, x1), dim=-1)
return x.flatten(-2)
def _apply_rotary_emb(
x: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
is_neox_style: bool,
) -> torch.Tensor:
"""
Args:
x: [num_tokens, num_heads, head_size]
cos: [num_tokens, head_size // 2]
sin: [num_tokens, head_size // 2]
is_neox_style: Whether to use the Neox-style or GPT-J-style rotary
positional embeddings.
"""
# cos = cos.unsqueeze(-2).to(x.dtype)
# sin = sin.unsqueeze(-2).to(x.dtype)
cos = cos.unsqueeze(-2)
sin = sin.unsqueeze(-2)
if is_neox_style:
x1, x2 = torch.chunk(x, 2, dim=-1)
else:
x1 = x[..., ::2]
x2 = x[..., 1::2]
o1 = (x1.float() * cos - x2.float() * sin).type_as(x)
o2 = (x2.float() * cos + x1.float() * sin).type_as(x)
if is_neox_style:
return torch.cat((o1, o2), dim=-1)
else:
return torch.stack((o1, o2), dim=-1).flatten(-2)
@CustomOp.register("rotary_embedding")
class RotaryEmbedding(CustomOp):
"""Original rotary positional embedding."""
def __init__(
self,
head_size: int,
rotary_dim: int,
max_position_embeddings: int,
base: int,
is_neox_style: bool,
dtype: torch.dtype,
) -> None:
super().__init__()
self.head_size = head_size
self.rotary_dim = rotary_dim
self.max_position_embeddings = max_position_embeddings
self.base = base
self.is_neox_style = is_neox_style
self.dtype = dtype
cache = self._compute_cos_sin_cache()
cache = cache.to(dtype)
self.cos_sin_cache: torch.Tensor
self.register_buffer("cos_sin_cache", cache, persistent=False)
def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor:
"""Compute the inverse frequency."""
# NOTE(woosuk): To exactly match the HF implementation, we need to
# use CPU to compute the cache and then move it to GPU. However, we
# create the cache on GPU for faster initialization. This may cause
# a slight numerical difference between the HF implementation and ours.
inv_freq = 1.0 / (base**(torch.arange(
0, self.rotary_dim, 2, dtype=torch.float) / self.rotary_dim))
return inv_freq
def _compute_cos_sin_cache(self) -> torch.Tensor:
"""Compute the cos and sin cache."""
inv_freq = self._compute_inv_freq(self.base)
t = torch.arange(self.max_position_embeddings, dtype=torch.float)
freqs = torch.einsum("i,j -> ij", t, inv_freq)
cos = freqs.cos()
sin = freqs.sin()
cache = torch.cat((cos, sin), dim=-1)
return cache
def forward_native(
self,
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
offsets: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""A PyTorch-native implementation of forward()."""
if offsets is not None:
positions = positions + offsets
positions = positions.flatten()
num_tokens = positions.shape[0]
cos_sin = self.cos_sin_cache.index_select(0, positions)
cos, sin = cos_sin.chunk(2, dim=-1)
query_shape = query.shape
query = query.view(num_tokens, -1, self.head_size)
query_rot = query[..., :self.rotary_dim]
query_pass = query[..., self.rotary_dim:]
query_rot = _apply_rotary_emb(query_rot, cos, sin, self.is_neox_style)
query = torch.cat((query_rot, query_pass), dim=-1).reshape(query_shape)
key_shape = key.shape
key = key.view(num_tokens, -1, self.head_size)
key_rot = key[..., :self.rotary_dim]
key_pass = key[..., self.rotary_dim:]
key_rot = _apply_rotary_emb(key_rot, cos, sin, self.is_neox_style)
key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape)
return query, key
def forward_cuda(
self,
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
offsets: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
from vllm import _custom_ops as ops
self.cos_sin_cache = self.cos_sin_cache.to(query.device,
dtype=query.dtype)
# ops.rotary_embedding()/batched_rotary_embedding()
# are in-place operations that update the query and key tensors.
if offsets is not None:
ops.batched_rotary_embedding(positions, query, key, self.head_size,
self.cos_sin_cache,
self.is_neox_style, self.rotary_dim,
offsets)
else:
ops.rotary_embedding(positions, query, key, self.head_size,
self.cos_sin_cache, self.is_neox_style)
return query, key
def forward_xpu(
self,
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
offsets: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
from vllm._ipex_ops import ipex_ops as ops
self.cos_sin_cache = self.cos_sin_cache.to(positions.device,
dtype=query.dtype)
# ops.rotary_embedding()/batched_rotary_embedding()
# are in-place operations that update the query and key tensors.
if offsets is not None:
ops.batched_rotary_embedding(positions, query, key, self.head_size,
self.cos_sin_cache,
self.is_neox_style, self.rotary_dim,
offsets)
else:
ops.rotary_embedding(positions, query, key, self.head_size,
self.cos_sin_cache, self.is_neox_style)
return query, key
def forward_hpu(
self,
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
offsets: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
from habana_frameworks.torch.hpex.kernels import (
RotaryPosEmbeddingMode, apply_rotary_pos_emb)
if offsets is not None:
offsets = offsets.view(positions.shape[0], -1)
positions = positions + offsets
positions = positions.flatten()
num_tokens = positions.shape[0]
cos_sin = self.cos_sin_cache.index_select(0, positions).view(
num_tokens, 1, -1)
cos, sin = cos_sin.chunk(2, dim=-1)
# HPU RoPE kernel requires hidden dimension for cos and sin to be equal
# to query hidden dimension, so the original tensors need to be
# expanded
# GPT-NeoX kernel requires position_ids = None, offset, mode = BLOCKWISE
# and expansion of cos/sin tensors via concatenation
# GPT-J kernel requires position_ids = None, offset = 0, mode = PAIRWISE
# and expansion of cos/sin tensors via repeat_interleave
rope_mode: RotaryPosEmbeddingMode
if self.is_neox_style:
rope_mode = RotaryPosEmbeddingMode.BLOCKWISE
cos = torch.cat((cos, cos), dim=-1)
sin = torch.cat((sin, sin), dim=-1)
else:
rope_mode = RotaryPosEmbeddingMode.PAIRWISE
sin = torch.repeat_interleave(sin,
2,
dim=-1,
output_size=cos_sin.shape[-1])
cos = torch.repeat_interleave(cos,
2,
dim=-1,
output_size=cos_sin.shape[-1])
query_shape = query.shape
query = query.view(num_tokens, -1, self.head_size)
query_rot = query[..., :self.rotary_dim]
query_pass = query[..., self.rotary_dim:]
query_rot = apply_rotary_pos_emb(query_rot, cos, sin, None, 0,
rope_mode)
query = torch.cat((query_rot, query_pass), dim=-1).reshape(query_shape)
key_shape = key.shape
key = key.view(num_tokens, -1, self.head_size)
key_rot = key[..., :self.rotary_dim]
key_pass = key[..., self.rotary_dim:]
key_rot = apply_rotary_pos_emb(key_rot, cos, sin, None, 0, rope_mode)
key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape)
return query, key
def extra_repr(self) -> str:
s = f"head_size={self.head_size}, rotary_dim={self.rotary_dim}"
s += f", max_position_embeddings={self.max_position_embeddings}"
s += f", base={self.base}, is_neox_style={self.is_neox_style}"
return s
def _to_tuple(x, dim=2):
if isinstance(x, int):
return (x,) * dim
elif len(x) == dim:
return x
else:
raise ValueError(f"Expected length {dim} or int, but got {x}")
def get_meshgrid_nd(start, *args, dim=2):
"""
Get n-D meshgrid with start, stop and num.
Args:
start (int or tuple): If len(args) == 0, start is num; If len(args) == 1, start is start, args[0] is stop,
step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num. For n-dim, start/stop/num
should be int or n-tuple. If n-tuple is provided, the meshgrid will be stacked following the dim order in
n-tuples.
*args: See above.
dim (int): Dimension of the meshgrid. Defaults to 2.
Returns:
grid (np.ndarray): [dim, ...]
"""
if len(args) == 0:
# start is grid_size
num = _to_tuple(start, dim=dim)
start = (0, ) * dim
stop = num
elif len(args) == 1:
# start is start, args[0] is stop, step is 1
start = _to_tuple(start, dim=dim)
stop = _to_tuple(args[0], dim=dim)
num = [stop[i] - start[i] for i in range(dim)]
elif len(args) == 2:
# start is start, args[0] is stop, args[1] is num
start = _to_tuple(start, dim=dim) # Left-Top eg: 12,0
stop = _to_tuple(args[0], dim=dim) # Right-Bottom eg: 20,32
num = _to_tuple(args[1], dim=dim) # Target Size eg: 32,124
else:
raise ValueError(f"len(args) should be 0, 1 or 2, but got {len(args)}")
# PyTorch implement of np.linspace(start[i], stop[i], num[i], endpoint=False)
axis_grid = []
for i in range(dim):
a, b, n = start[i], stop[i], num[i]
g = torch.linspace(a, b, n + 1, dtype=torch.float32)[:n]
axis_grid.append(g)
grid = torch.meshgrid(*axis_grid, indexing="ij") # dim x [W, H, D]
grid = torch.stack(grid, dim=0) # [dim, W, H, D]
return grid
def get_1d_rotary_pos_embed(
dim: int,
pos: Union[torch.FloatTensor, int],
theta: float = 10000.0,
theta_rescale_factor: float = 1.0,
interpolation_factor: float = 1.0,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
(Note: `cis` means `cos + i * sin`, where i is the imaginary unit.)
This function calculates a frequency tensor with complex exponential using the given dimension 'dim'
and the end index 'end'. The 'theta' parameter scales the frequencies.
Args:
dim (int): Dimension of the frequency tensor.
pos (int or torch.FloatTensor): Position indices for the frequency tensor. [S] or scalar
theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.
theta_rescale_factor (float, optional): Rescale factor for theta. Defaults to 1.0.
interpolation_factor (float, optional): Factor to scale positions. Defaults to 1.0.
Returns:
freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately. [S, D]
"""
if isinstance(pos, int):
pos = torch.arange(pos).float()
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
# has some connection to NTK literature
if theta_rescale_factor != 1.0:
theta *= theta_rescale_factor**(dim / (dim - 2))
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].to(torch.float64) / dim)) # [D/2]
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
freqs_cos = freqs.cos() # [S, D/2]
freqs_sin = freqs.sin() # [S, D/2]
return freqs_cos, freqs_sin
def get_nd_rotary_pos_embed(
rope_dim_list,
start,
*args,
theta=10000.0,
theta_rescale_factor: Union[float, List[float]] = 1.0,
interpolation_factor: Union[float, List[float]] = 1.0,
shard_dim: int = 0,
sp_rank: int = 0,
sp_world_size: int = 1,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
Supports sequence parallelism by allowing sharding of a specific dimension.
Args:
rope_dim_list (list of int): Dimension of each rope. len(rope_dim_list) should equal to n.
sum(rope_dim_list) should equal to head_dim of attention layer.
start (int | tuple of int | list of int): If len(args) == 0, start is num; If len(args) == 1, start is start,
args[0] is stop, step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num.
*args: See above.
theta (float): Scaling factor for frequency computation. Defaults to 10000.0.
theta_rescale_factor (float): Rescale factor for theta. Defaults to 1.0.
interpolation_factor (float): Factor to scale positions. Defaults to 1.0.
shard_dim (int): Which dimension to shard for sequence parallelism. Defaults to 0.
sp_rank (int): Rank in the sequence parallel group. Defaults to 0.
sp_world_size (int): World size of the sequence parallel group. Defaults to 1.
Returns:
Tuple[torch.Tensor, torch.Tensor]: (cos, sin) tensors of shape [HW, D/2]
"""
# Get the full grid
full_grid = get_meshgrid_nd(start, *args, dim=len(rope_dim_list)) # [3, W, H, D] / [2, W, H]
# Shard the grid if using sequence parallelism (sp_world_size > 1)
assert shard_dim < len(rope_dim_list), f"shard_dim {shard_dim} must be less than number of dimensions {len(rope_dim_list)}"
if sp_world_size > 1:
# Get the shape of the full grid
grid_shape = list(full_grid.shape[1:])
# Ensure the dimension to shard is divisible by sp_world_size
assert grid_shape[shard_dim] % sp_world_size == 0, (
f"Dimension {shard_dim} with size {grid_shape[shard_dim]} is not divisible "
f"by sequence parallel world size {sp_world_size}"
)
# Compute the start and end indices for this rank's shard
shard_size = grid_shape[shard_dim] // sp_world_size
start_idx = sp_rank * shard_size
end_idx = (sp_rank + 1) * shard_size
# Create slicing indices for each dimension
slice_indices = [slice(None) for _ in range(len(grid_shape))]
slice_indices[shard_dim] = slice(start_idx, end_idx)
# Shard the grid
# Update grid shape for the sharded dimension
grid_shape[shard_dim] = grid_shape[shard_dim] // sp_world_size
grid = torch.empty((len(rope_dim_list),) + tuple(grid_shape), dtype=full_grid.dtype)
for i in range(len(rope_dim_list)):
grid[i] = full_grid[i][tuple(slice_indices)]
else:
grid = full_grid
if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
assert len(theta_rescale_factor) == len(
rope_dim_list), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
assert len(interpolation_factor) == len(
rope_dim_list), "len(interpolation_factor) should equal to len(rope_dim_list)"
# use 1/ndim of dimensions to encode grid_axis
embs = []
for i in range(len(rope_dim_list)):
emb = get_1d_rotary_pos_embed(
rope_dim_list[i],
grid[i].reshape(-1),
theta,
theta_rescale_factor=theta_rescale_factor[i],
interpolation_factor=interpolation_factor[i],
) # 2 x [WHD, rope_dim_list[i]]
embs.append(emb)
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D/2)
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D/2)
return cos, sin
def get_rotary_pos_embed(
rope_sizes,
hidden_size,
heads_num,
rope_dim_list,
rope_theta,
theta_rescale_factor=1.0,
interpolation_factor=1.0,
shard_dim: int = 0,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Generate rotary positional embeddings for the given sizes.
Args:
rope_sizes: Tuple of dimensions (t, h, w)
hidden_size: Hidden dimension size
heads_num: Number of attention heads
rope_dim_list: List of dimensions for each axis, or None
rope_theta: Base for frequency calculations
theta_rescale_factor: Rescale factor for theta. Defaults to 1.0
interpolation_factor: Factor to scale positions. Defaults to 1.0
shard_dim: Which dimension to shard for sequence parallelism. Defaults to 0.
Returns:
Tuple of (cos, sin) tensors for rotary embeddings
"""
target_ndim = 3
head_dim = hidden_size // heads_num
if rope_dim_list is None:
rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)]
assert sum(rope_dim_list) == head_dim, "sum(rope_dim_list) should equal to head_dim of attention layer"
# Get SP info
sp_group = get_sp_group()
sp_rank = sp_group.rank_in_group
sp_world_size = sp_group.world_size
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
rope_dim_list,
rope_sizes,
theta=rope_theta,
theta_rescale_factor=theta_rescale_factor,
interpolation_factor=interpolation_factor,
shard_dim=shard_dim,
sp_rank=sp_rank,
sp_world_size=sp_world_size
)
return freqs_cos, freqs_sin
_ROPE_DICT: Dict[Tuple, RotaryEmbedding] = {}
def get_rope(
head_size: int,
rotary_dim: int,
max_position: int,
base: int,
is_neox_style: bool = True,
rope_scaling: Optional[Dict[str, Any]] = None,
dtype: Optional[torch.dtype] = None,
partial_rotary_factor: float = 1.0,
) -> RotaryEmbedding:
if dtype is None:
dtype = torch.get_default_dtype()
if rope_scaling is not None:
# Transforms every value that is a list into a tuple for caching calls
rope_scaling_tuple = {
k: tuple(v) if isinstance(v, list) else v
for k, v in rope_scaling.items()
}
rope_scaling_args = tuple(rope_scaling_tuple.items())
else:
rope_scaling_args = None
if partial_rotary_factor < 1.0:
rotary_dim = int(rotary_dim * partial_rotary_factor)
key = (head_size, rotary_dim, max_position, base, is_neox_style,
rope_scaling_args, dtype)
if key in _ROPE_DICT:
return _ROPE_DICT[key]
if rope_scaling is None:
rotary_emb = RotaryEmbedding(head_size, rotary_dim, max_position, base,
is_neox_style, dtype)
else:
raise ValueError(f"Unknown RoPE scaling {rope_scaling}")
_ROPE_DICT[key] = rotary_emb
return rotary_emb
-58
View File
@@ -1,58 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Utility methods for model layers."""
from typing import Tuple
import torch
def get_token_bin_counts_and_mask(
tokens: torch.Tensor,
vocab_size: int,
num_seqs: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
# Compute the bin counts for the tokens.
# vocab_size + 1 for padding.
bin_counts = torch.zeros((num_seqs, vocab_size + 1),
dtype=torch.long,
device=tokens.device)
bin_counts.scatter_add_(1, tokens, torch.ones_like(tokens))
bin_counts = bin_counts[:, :vocab_size]
mask = bin_counts > 0
return bin_counts, mask
def apply_penalties(logits: torch.Tensor, prompt_tokens_tensor: torch.Tensor,
output_tokens_tensor: torch.Tensor,
presence_penalties: torch.Tensor,
frequency_penalties: torch.Tensor,
repetition_penalties: torch.Tensor) -> torch.Tensor:
"""
Applies penalties in place to the logits tensor
logits : The input logits tensor of shape [num_seqs, vocab_size]
prompt_tokens_tensor: A tensor containing the prompt tokens. The prompts
are padded to the maximum prompt length within the batch using
`vocab_size` as the padding value. The value `vocab_size` is used
for padding because it does not correspond to any valid token ID
in the vocabulary.
output_tokens_tensor: The output tokens tensor.
presence_penalties: The presence penalties of shape (num_seqs, )
frequency_penalties: The frequency penalties of shape (num_seqs, )
repetition_penalties: The repetition penalties of shape (num_seqs, )
"""
num_seqs, vocab_size = logits.shape
_, prompt_mask = get_token_bin_counts_and_mask(prompt_tokens_tensor,
vocab_size, num_seqs)
output_bin_counts, output_mask = get_token_bin_counts_and_mask(
output_tokens_tensor, vocab_size, num_seqs)
repetition_penalties = repetition_penalties.unsqueeze(dim=1).repeat(
1, vocab_size)
logits[logits > 0] /= torch.where(prompt_mask | output_mask,
repetition_penalties, 1.0)[logits > 0]
logits[logits <= 0] *= torch.where(prompt_mask | output_mask,
repetition_penalties, 1.0)[logits <= 0]
# We follow the definition in OpenAI API.
# Refer to https://platform.openai.com/docs/api-reference/parameter-details
logits -= frequency_penalties.unsqueeze(dim=1) * output_bin_counts
logits -= presence_penalties.unsqueeze(dim=1) * output_mask
return logits
-165
View File
@@ -1,165 +0,0 @@
import torch
import torch.nn as nn
import math
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.layers.linear import ReplicatedLinear
from typing import Optional
from fastvideo.v1.layers.mlp import MLP
class PatchEmbed(nn.Module):
"""2D Image to Patch Embedding
Image to Patch Embedding using Conv2d
A convolution based approach to patchifying a 2D image w/ embedding projection.
Based on the impl in https://github.com/google-research/vision_transformer
Hacked together by / Copyright 2020 Ross Wightman
Remove the _assert function in forward function to be compatible with multi-resolution images.
"""
def __init__(
self,
patch_size=16,
in_chans=3,
embed_dim=768,
norm_layer=None,
flatten=True,
bias=True,
dtype=None
):
super().__init__()
# Convert patch_size to 2-tuple
if isinstance(patch_size, (list, tuple)):
if len(patch_size) == 1:
patch_size = (patch_size[0], patch_size[0])
else:
patch_size = (patch_size, patch_size)
self.patch_size = patch_size
self.flatten = flatten
self.proj = nn.Conv3d(
in_chans,
embed_dim,
kernel_size=patch_size,
stride=patch_size,
bias=bias,
dtype=dtype
)
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
def forward(self, x):
x = self.proj(x)
if self.flatten:
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
x = self.norm(x)
return x
class TimestepEmbedder(nn.Module):
"""
Embeds scalar timesteps into vector representations.
"""
def __init__(
self,
hidden_size,
act_layer="silu",
frequency_embedding_size=256,
max_period=10000,
dtype=None,
):
super().__init__()
self.frequency_embedding_size = frequency_embedding_size
self.max_period = max_period
self.mlp = MLP(
frequency_embedding_size,
hidden_size,
hidden_size,
act_type=act_layer,
dtype=dtype
)
def forward(self, t):
t_freq = timestep_embedding(t, self.frequency_embedding_size, self.max_period).float()
# t_freq = t_freq.to(self.mlp.fc_in.weight.dtype)
t_emb = self.mlp(t_freq)
return t_emb
def timestep_embedding(t, dim, max_period=10000):
"""
Create sinusoidal timestep embeddings.
Args:
t: Tensor of shape [B] with timesteps
dim: Embedding dimension
max_period: Controls the minimum frequency of the embeddings
Returns:
Tensor of shape [B, dim] with embeddings
"""
half = dim // 2
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float64) / half).to(device=t.device)
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
class ModulateProjection(nn.Module):
"""Modulation layer for DiT blocks."""
def __init__(
self,
hidden_size: int,
factor: int = 2,
act_layer: str = "silu",
dtype: Optional[torch.dtype] = None,
):
super().__init__()
self.factor = factor
self.hidden_size = hidden_size
self.linear = ReplicatedLinear(
hidden_size,
hidden_size * factor,
bias=True,
params_dtype=dtype
)
self.act = get_act_fn(act_layer)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.act(x)
x, _ = self.linear(x)
return x
def unpatchify(x, t, h, w, patch_size, channels):
"""
Convert patched representation back to image space.
Args:
x: Tensor of shape [B, T*H*W, C*P_t*P_h*P_w]
t, h, w: Temporal and spatial dimensions
Returns:
Unpatchified tensor of shape [B, C, T*P_t, H*P_h, W*P_w]
"""
assert x.ndim == 3, f"x.ndim: {x.ndim}"
assert len(patch_size) == 3, f"patch_size: {patch_size}"
assert t * h * w == x.shape[1], f"t * h * w: {t * h * w}, x.shape[1]: {x.shape[1]}"
c = channels
pt, ph, pw = patch_size
x = x.reshape(shape=(x.shape[0], t, h, w, c, pt, ph, pw))
x = torch.einsum("nthwcopq->nctohpwq", x)
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
return imgs
@@ -1,484 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from typing import List, Optional, Sequence, Tuple
import torch
import torch.nn.functional as F
from torch.nn.parameter import Parameter, UninitializedParameter
from fastvideo.v1.distributed import (divide, get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce)
from vllm.model_executor.layers.quantization.base_config import (
QuantizationConfig, QuantizeMethodBase, method_has_implemented_embedding)
from fastvideo.v1.models.parameter import BasevLLMParameter
from fastvideo.v1.models.utils import set_weight_attrs
from fastvideo.v1.platforms import current_platform
DEFAULT_VOCAB_PADDING_SIZE = 64
class UnquantizedEmbeddingMethod(QuantizeMethodBase):
"""Unquantized method for embeddings."""
def create_weights(self, layer: torch.nn.Module,
input_size_per_partition: int,
output_partition_sizes: List[int], input_size: int,
output_size: int, params_dtype: torch.dtype,
**extra_weight_attrs):
"""Create weights for embedding layer."""
weight = Parameter(torch.empty(sum(output_partition_sizes),
input_size_per_partition,
dtype=params_dtype),
requires_grad=False)
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
layer.register_parameter("weight", weight)
set_weight_attrs(weight, extra_weight_attrs)
def apply(self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None) -> torch.Tensor:
return F.linear(x, layer.weight, bias)
def embedding(self, layer: torch.nn.Module,
input_: torch.Tensor) -> torch.Tensor:
return F.embedding(input_, layer.weight)
def pad_vocab_size(vocab_size: int,
pad_to: int = DEFAULT_VOCAB_PADDING_SIZE) -> int:
"""Pad the vocab size to the given value."""
return ((vocab_size + pad_to - 1) // pad_to) * pad_to
def vocab_range_from_per_partition_vocab_size(
per_partition_vocab_size: int,
rank: int,
offset: int = 0) -> Sequence[int]:
index_f = rank * per_partition_vocab_size
index_l = index_f + per_partition_vocab_size
return index_f + offset, index_l + offset
def vocab_range_from_global_vocab_size(global_vocab_size: int,
rank: int,
world_size: int,
offset: int = 0) -> Sequence[int]:
per_partition_vocab_size = divide(global_vocab_size, world_size)
return vocab_range_from_per_partition_vocab_size(per_partition_vocab_size,
rank,
offset=offset)
@dataclass
class VocabParallelEmbeddingShardIndices:
"""Indices for a shard of a vocab parallel embedding."""
padded_org_vocab_start_index: int
padded_org_vocab_end_index: int
padded_added_vocab_start_index: int
padded_added_vocab_end_index: int
org_vocab_start_index: int
org_vocab_end_index: int
added_vocab_start_index: int
added_vocab_end_index: int
@property
def num_org_elements(self) -> int:
return self.org_vocab_end_index - self.org_vocab_start_index
@property
def num_added_elements(self) -> int:
return self.added_vocab_end_index - self.added_vocab_start_index
@property
def num_org_elements_padded(self) -> int:
return (self.padded_org_vocab_end_index -
self.padded_org_vocab_start_index)
@property
def num_added_elements_padded(self) -> int:
return (self.padded_added_vocab_end_index -
self.padded_added_vocab_start_index)
@property
def num_org_vocab_padding(self) -> int:
return self.num_org_elements_padded - self.num_org_elements
@property
def num_added_vocab_padding(self) -> int:
return self.num_added_elements_padded - self.num_added_elements
@property
def num_elements_padded(self) -> int:
return self.num_org_elements_padded + self.num_added_elements_padded
def __post_init__(self):
# sanity checks
assert (self.padded_org_vocab_start_index
<= self.padded_org_vocab_end_index)
assert (self.padded_added_vocab_start_index
<= self.padded_added_vocab_end_index)
assert self.org_vocab_start_index <= self.org_vocab_end_index
assert self.added_vocab_start_index <= self.added_vocab_end_index
assert self.org_vocab_start_index <= self.padded_org_vocab_start_index
assert (self.added_vocab_start_index
<= self.padded_added_vocab_start_index)
assert self.org_vocab_end_index <= self.padded_org_vocab_end_index
assert self.added_vocab_end_index <= self.padded_added_vocab_end_index
assert self.num_org_elements <= self.num_org_elements_padded
assert self.num_added_elements <= self.num_added_elements_padded
@torch.compile(dynamic=True, backend=current_platform.simple_compile_backend)
def get_masked_input_and_mask(
input_: torch.Tensor, org_vocab_start_index: int,
org_vocab_end_index: int, num_org_vocab_padding: int,
added_vocab_start_index: int,
added_vocab_end_index: int) -> Tuple[torch.Tensor, torch.Tensor]:
# torch.compile will fuse all of the pointwise ops below
# into a single kernel, making it very fast
org_vocab_mask = (input_ >= org_vocab_start_index) & (
input_ < org_vocab_end_index)
added_vocab_mask = (input_ >= added_vocab_start_index) & (
input_ < added_vocab_end_index)
added_offset = added_vocab_start_index - (
org_vocab_end_index - org_vocab_start_index) - num_org_vocab_padding
valid_offset = (org_vocab_start_index *
org_vocab_mask) + (added_offset * added_vocab_mask)
vocab_mask = org_vocab_mask | added_vocab_mask
input_ = vocab_mask * (input_ - valid_offset)
return input_, ~vocab_mask
class VocabParallelEmbedding(torch.nn.Module):
"""Embedding parallelized in the vocabulary dimension.
Adapted from torch.nn.Embedding, note that we pad the vocabulary size to
make sure it is divisible by the number of model parallel GPUs.
In order to support various loading methods, we ensure that LoRA-added
embeddings are always at the end of TP-sharded tensors. In other words,
we shard base embeddings and LoRA embeddings separately (both padded),
and place them in the same tensor.
In this example, we will have the original vocab size = 1010,
added vocab size = 16 and padding to 64. Therefore, the total
vocab size with padding will be 1088 (because we first pad 1010 to
1024, add 16, and then pad to 1088).
Therefore, the tensor format looks like the following:
TP1, rank 0 (no sharding):
|< --------BASE-------- >|< -BASE PADDING-- >|< -----LORA------ >|< -LORA PADDING-- >|
corresponding token_id: | 0 | 1 | ... | 1009 | -1 | ... | -1 | 1010 | ... | 1015 | -1 | ... | -1 |
index: | 0 | 1 | ... | 1009 | 1010 | ... | 1023 | 1024 | ... | 1039 | 1040 | ... | 1087 |
TP2, rank 0:
|< --------------------BASE--------------------- >|< -----LORA------ >|< -LORA PADDING- >|
corresponding token_id: | 0 | 1 | 2 | ... | 497 | 498 | ... | 511 | 1000 | ... | 1015 | -1 | ... | -1 |
index: | 0 | 1 | 2 | ... | 497 | 498 | ... | 511 | 512 | ... | 527 | 520 | ... | 543 |
TP2, rank 1:
|< -----------BASE----------- >|< -BASE PADDING- >|< -----------LORA PADDING----------- >|
corresponding token_id: | 512 | 513 | 514 | ... | 1009 | -1 | ... | -1 | -1 | ... | -1 | -1 | ... | -1 |
index: | 0 | 1 | 2 | ... | 497 | 498 | ... | 511 | 512 | ... | 519 | 520 | ... | 543 |
Args:
num_embeddings: vocabulary size.
embedding_dim: size of hidden state.
params_dtype: type of the parameters.
org_num_embeddings: original vocabulary size (without LoRA).
padding_size: padding size for the vocabulary.
quant_config: quant config for the layer
prefix: full name of the layer in the state dict
""" # noqa: E501
def __init__(self,
num_embeddings: int,
embedding_dim: int,
params_dtype: Optional[torch.dtype] = None,
org_num_embeddings: Optional[int] = None,
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
super().__init__()
# Keep the input dimensions.
tp_rank = get_tensor_model_parallel_rank()
self.tp_size = get_tensor_model_parallel_world_size()
self.num_embeddings = num_embeddings
self.padding_size = padding_size
self.org_vocab_size = org_num_embeddings or num_embeddings
num_added_embeddings = num_embeddings - self.org_vocab_size
self.org_vocab_size_padded = pad_vocab_size(self.org_vocab_size,
self.padding_size)
self.num_embeddings_padded = pad_vocab_size(
self.org_vocab_size_padded + num_added_embeddings,
self.padding_size)
assert self.org_vocab_size_padded <= self.num_embeddings_padded
self.shard_indices = self._get_indices(self.num_embeddings_padded,
self.org_vocab_size_padded,
self.num_embeddings,
self.org_vocab_size, tp_rank,
self.tp_size)
self.embedding_dim = embedding_dim
quant_method = None
if quant_config is not None:
quant_method = quant_config.get_quant_method(self, prefix=prefix)
if quant_method is None:
quant_method = UnquantizedEmbeddingMethod()
# If we are making an embedding layer, then our quantization linear
# method must implement the embedding operation. If we are another
# layer type like ParallelLMHead, this is not important.
is_embedding_layer = type(self.__class__) is VocabParallelEmbedding
quant_method_implements_embedding = method_has_implemented_embedding(
type(quant_method))
if is_embedding_layer and not quant_method_implements_embedding:
raise NotImplementedError(
f"The class {type(quant_method).__name__} must implement "
"the 'embedding' method, see UnquantizedEmbeddingMethod.")
self.quant_method: QuantizeMethodBase = quant_method
if params_dtype is None:
params_dtype = torch.get_default_dtype()
# Divide the weight matrix along the vocaburaly dimension.
self.num_added_embeddings = self.num_embeddings - self.org_vocab_size
self.num_embeddings_per_partition = divide(self.num_embeddings_padded,
self.tp_size)
assert (self.shard_indices.num_elements_padded ==
self.num_embeddings_per_partition)
self.num_org_embeddings_per_partition = (
self.shard_indices.org_vocab_end_index -
self.shard_indices.org_vocab_start_index)
self.num_added_embeddings_per_partition = (
self.shard_indices.added_vocab_end_index -
self.shard_indices.added_vocab_start_index)
self.quant_method.create_weights(self,
self.embedding_dim,
[self.num_embeddings_per_partition],
self.embedding_dim,
self.num_embeddings_padded,
params_dtype=params_dtype,
weight_loader=self.weight_loader)
@classmethod
def _get_indices(cls, vocab_size_padded: int, org_vocab_size_padded: int,
vocab_size: int, org_vocab_size: int, tp_rank: int,
tp_size: int) -> VocabParallelEmbeddingShardIndices:
"""Get start and end indices for vocab parallel embedding, following the
layout outlined in the class docstring, based on the given tp_rank and
tp_size."""
num_added_embeddings_padded = vocab_size_padded - org_vocab_size_padded
padded_org_vocab_start_index, padded_org_vocab_end_index = (
vocab_range_from_global_vocab_size(org_vocab_size_padded, tp_rank,
tp_size))
padded_added_vocab_start_index, padded_added_vocab_end_index = (
vocab_range_from_global_vocab_size(num_added_embeddings_padded,
tp_rank,
tp_size,
offset=org_vocab_size))
# remove padding
org_vocab_start_index = min(padded_org_vocab_start_index,
org_vocab_size)
org_vocab_end_index = min(padded_org_vocab_end_index, org_vocab_size)
added_vocab_start_index = min(padded_added_vocab_start_index,
vocab_size)
added_vocab_end_index = min(padded_added_vocab_end_index, vocab_size)
return VocabParallelEmbeddingShardIndices(
padded_org_vocab_start_index, padded_org_vocab_end_index,
padded_added_vocab_start_index, padded_added_vocab_end_index,
org_vocab_start_index, org_vocab_end_index,
added_vocab_start_index, added_vocab_end_index)
def get_sharded_to_full_mapping(self) -> Optional[List[int]]:
"""Get a mapping that can be used to reindex the gathered
logits for sampling.
During sampling, we gather logits from all ranks. The relationship
of index->token_id will follow the same format as outlined in the class
docstring. However, after the gather, we want to reindex the final
logits tensor to map index->token_id one-to-one (the index is always
equal the token_id it corresponds to). The indices returned by this
method allow us to do that.
"""
if self.tp_size < 2:
return None
base_embeddings: List[int] = []
added_embeddings: List[int] = []
padding: List[int] = []
for tp_rank in range(self.tp_size):
shard_indices = self._get_indices(self.num_embeddings_padded,
self.org_vocab_size_padded,
self.num_embeddings,
self.org_vocab_size, tp_rank,
self.tp_size)
range_start = self.num_embeddings_per_partition * tp_rank
range_end = self.num_embeddings_per_partition * (tp_rank + 1)
base_embeddings.extend(
range(range_start,
range_start + shard_indices.num_org_elements))
padding.extend(
range(range_start + shard_indices.num_org_elements,
range_start + shard_indices.num_org_elements_padded))
added_embeddings.extend(
range(
range_start + shard_indices.num_org_elements_padded,
range_start + shard_indices.num_org_elements_padded +
shard_indices.num_added_elements))
padding.extend(
range(
range_start + shard_indices.num_org_elements_padded +
shard_indices.num_added_elements,
range_start + shard_indices.num_org_elements_padded +
shard_indices.num_added_elements_padded))
assert (range_start + shard_indices.num_org_elements_padded +
shard_indices.num_added_elements_padded == range_end)
ret = base_embeddings + added_embeddings + padding
assert len(ret) == self.num_embeddings_padded
return ret
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
output_dim = getattr(param, "output_dim", None)
packed_dim = getattr(param, "packed_dim", None)
# If the parameter is a gguf weight, then load it directly.
if getattr(param, "is_gguf_weight_type", None):
param.data.copy_(loaded_weight)
param.weight_type = loaded_weight.item()
return
elif isinstance(param, UninitializedParameter):
shape = list(loaded_weight.shape)
if output_dim is not None:
shape[output_dim] = self.num_embeddings_per_partition
param.materialize(tuple(shape), dtype=loaded_weight.dtype)
# If parameter does not have output dim, then it should
# be copied onto all gpus (e.g. g_idx for act_order gptq).
if output_dim is None:
assert param.data.shape == loaded_weight.shape
param.data.copy_(loaded_weight)
return
# Shard indexes for loading the weight
start_idx = self.shard_indices.org_vocab_start_index
shard_size = self.shard_indices.org_vocab_end_index - start_idx
# If param packed on the same dim we are sharding on, then
# need to adjust offsets of loaded weight by pack_factor.
if packed_dim is not None and packed_dim == output_dim:
packed_factor = param.packed_factor if isinstance(
param, BasevLLMParameter) else param.pack_factor
assert loaded_weight.shape[output_dim] == (self.org_vocab_size //
param.packed_factor)
start_idx = start_idx // packed_factor
shard_size = shard_size // packed_factor
else:
assert loaded_weight.shape[output_dim] == self.org_vocab_size
# Copy the data. Select chunk corresponding to current shard.
loaded_weight = loaded_weight.narrow(output_dim, start_idx, shard_size)
if current_platform.is_hpu():
# FIXME(kzawora): Weight copy with slicing bugs out on Gaudi here,
# so we're using a workaround. Remove this when fixed in
# HPU PT bridge.
padded_weight = torch.cat([
loaded_weight,
torch.zeros(param.shape[0] - loaded_weight.shape[0],
*loaded_weight.shape[1:])
])
param.data.copy_(padded_weight)
else:
param[:loaded_weight.shape[0]].data.copy_(loaded_weight)
param[loaded_weight.shape[0]:].data.fill_(0)
def forward(self, input_):
if self.tp_size > 1:
# Build the mask.
masked_input, input_mask = get_masked_input_and_mask(
input_, self.shard_indices.org_vocab_start_index,
self.shard_indices.org_vocab_end_index,
self.shard_indices.num_org_vocab_padding,
self.shard_indices.added_vocab_start_index,
self.shard_indices.added_vocab_end_index)
else:
masked_input = input_
# Get the embeddings.
output_parallel = self.quant_method.embedding(self,
masked_input.long())
# Mask the output embedding.
if self.tp_size > 1:
output_parallel.masked_fill_(input_mask.unsqueeze(-1), 0)
# Reduce across all the model parallel GPUs.
output = tensor_model_parallel_all_reduce(output_parallel)
return output
def extra_repr(self) -> str:
s = f"num_embeddings={self.num_embeddings_per_partition}"
s += f", embedding_dim={self.embedding_dim}"
s += f", org_vocab_size={self.org_vocab_size}"
s += f', num_embeddings_padded={self.num_embeddings_padded}'
s += f', tp_size={self.tp_size}'
return s
class ParallelLMHead(VocabParallelEmbedding):
"""Parallelized LM head.
Output logits weight matrices used in the Sampler. The weight and bias
tensors are padded to make sure they are divisible by the number of
model parallel GPUs.
Args:
num_embeddings: vocabulary size.
embedding_dim: size of hidden state.
bias: whether to use bias.
params_dtype: type of the parameters.
org_num_embeddings: original vocabulary size (without LoRA).
padding_size: padding size for the vocabulary.
"""
def __init__(self,
num_embeddings: int,
embedding_dim: int,
bias: bool = False,
params_dtype: Optional[torch.dtype] = None,
org_num_embeddings: Optional[int] = None,
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
super().__init__(num_embeddings, embedding_dim, params_dtype,
org_num_embeddings, padding_size, quant_config,
prefix)
self.quant_config = quant_config
if bias:
self.bias = Parameter(
torch.empty(self.num_embeddings_per_partition,
dtype=params_dtype))
set_weight_attrs(self.bias, {
"output_dim": 0,
"weight_loader": self.weight_loader,
})
else:
self.register_parameter("bias", None)
def tie_weights(self, embed_tokens: VocabParallelEmbedding):
"""Tie the weights with word embeddings."""
# GGUF quantized embed_tokens.
if self.quant_config and self.quant_config.get_name() == "gguf":
return embed_tokens
else:
self.weight = embed_tokens.weight
return self
def forward(self, input_):
del input_
raise RuntimeError("LMHead's weights should be used in the sampler.")
-219
View File
@@ -1,219 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from vllm
# https://github.com/vllm-project/vllm/blob/main/vllm/logger.py
# Copyright 2023 The vLLM Authors.
# Copyright 2023 The FastVideo Authors.
"""Logging configuration for fastvideo.v1."""
import datetime
import json
import logging
import os
import sys
from functools import lru_cache, partial
from logging import Logger
from logging.config import dictConfig
from os import path
from types import MethodType
from typing import Any, Optional, cast
import fastvideo.v1.envs as envs
FASTVIDEO_CONFIGURE_LOGGING = envs.FASTVIDEO_CONFIGURE_LOGGING
FASTVIDEO_LOGGING_CONFIG_PATH = envs.FASTVIDEO_LOGGING_CONFIG_PATH
FASTVIDEO_LOGGING_LEVEL = envs.FASTVIDEO_LOGGING_LEVEL
FASTVIDEO_LOGGING_PREFIX = envs.FASTVIDEO_LOGGING_PREFIX
_FORMAT = (f"{FASTVIDEO_LOGGING_PREFIX}%(levelname)s %(asctime)s "
"[%(filename)s:%(lineno)d] %(message)s")
_DATE_FORMAT = "%m-%d %H:%M:%S"
DEFAULT_LOGGING_CONFIG = {
"formatters": {
"fastvideo": {
"class": "fastvideo.v1.logging_utils.NewLineFormatter",
"datefmt": _DATE_FORMAT,
"format": _FORMAT,
},
},
"handlers": {
"fastvideo": {
"class": "logging.StreamHandler",
"formatter": "fastvideo",
"level": FASTVIDEO_LOGGING_LEVEL,
"stream": "ext://sys.stdout",
},
},
"loggers": {
"fastvideo": {
"handlers": ["fastvideo"],
"level": "DEBUG",
"propagate": False,
},
},
"root": {
"handlers": ["fastvideo"],
"level": "DEBUG",
},
"version": 1,
"disable_existing_loggers": False
}
@lru_cache
def _print_info_once(logger: Logger, msg: str) -> None:
# Set the stacklevel to 2 to print the original caller's line info
logger.info(msg, stacklevel=2)
@lru_cache
def _print_warning_once(logger: Logger, msg: str) -> None:
# Set the stacklevel to 2 to print the original caller's line info
logger.warning(msg, stacklevel=2)
class _FastvideoLogger(Logger):
"""
Note:
This class is just to provide type information.
We actually patch the methods directly on the :class:`logging.Logger`
instance to avoid conflicting with other libraries such as
`intel_extension_for_pytorch.utils._logger`.
"""
def info_once(self, msg: str) -> None:
"""
As :meth:`info`, but subsequent calls with the same message
are silently dropped.
"""
_print_info_once(self, msg)
def warning_once(self, msg: str) -> None:
"""
As :meth:`warning`, but subsequent calls with the same message
are silently dropped.
"""
_print_warning_once(self, msg)
def _configure_fastvideo_root_logger() -> None:
logging_config = dict[str, Any]()
if not FASTVIDEO_CONFIGURE_LOGGING and FASTVIDEO_LOGGING_CONFIG_PATH:
raise RuntimeError(
"FASTVIDEO_CONFIGURE_LOGGING evaluated to false, but "
"FASTVIDEO_LOGGING_CONFIG_PATH was given. FASTVIDEO_LOGGING_CONFIG_PATH "
"implies FASTVIDEO_CONFIGURE_LOGGING. Please enable "
"FASTVIDEO_CONFIGURE_LOGGING or unset FASTVIDEO_LOGGING_CONFIG_PATH.")
if FASTVIDEO_CONFIGURE_LOGGING:
logging_config = DEFAULT_LOGGING_CONFIG
if FASTVIDEO_LOGGING_CONFIG_PATH:
if not path.exists(FASTVIDEO_LOGGING_CONFIG_PATH):
raise RuntimeError(
"Could not load logging config. File does not exist: %s",
FASTVIDEO_LOGGING_CONFIG_PATH)
with open(FASTVIDEO_LOGGING_CONFIG_PATH, encoding="utf-8") as file:
custom_config = json.loads(file.read())
if not isinstance(custom_config, dict):
raise ValueError("Invalid logging config. Expected Dict, got %s.",
type(custom_config).__name__)
logging_config = custom_config
for formatter in logging_config.get("formatters", {}).values():
# This provides backwards compatibility after #10134.
if formatter.get("class") == "fastvideo.v1.logging.NewLineFormatter":
formatter["class"] = "fastvideo.v1.logging_utils.NewLineFormatter"
if logging_config:
dictConfig(logging_config)
# TODO: add rank_zero_only log
def init_logger(name: str) -> _FastvideoLogger:
"""The main purpose of this function is to ensure that loggers are
retrieved in such a way that we can be sure the root fastvideo logger has
already been configured."""
logger = logging.getLogger(name)
methods_to_patch = {
"info_once": _print_info_once,
"warning_once": _print_warning_once,
}
for method_name, method in methods_to_patch.items():
setattr(logger, method_name, MethodType(method, logger))
return cast(_FastvideoLogger, logger)
# The root logger is initialized when the module is imported.
# This is thread-safe as the module is only imported once,
# guaranteed by the Python GIL.
_configure_fastvideo_root_logger()
logger = init_logger(__name__)
def _trace_calls(log_path, root_dir, frame, event, arg=None):
if event in ['call', 'return']:
# Extract the filename, line number, function name, and the code object
filename = frame.f_code.co_filename
lineno = frame.f_lineno
func_name = frame.f_code.co_name
if not filename.startswith(root_dir):
# only log the functions in the fastvideo root_dir
return
# Log every function call or return
try:
last_frame = frame.f_back
if last_frame is not None:
last_filename = last_frame.f_code.co_filename
last_lineno = last_frame.f_lineno
last_func_name = last_frame.f_code.co_name
else:
# initial frame
last_filename = ""
last_lineno = 0
last_func_name = ""
with open(log_path, 'a') as f:
ts = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S.%f")
if event == 'call':
f.write(f"{ts} Call to"
f" {func_name} in {filename}:{lineno}"
f" from {last_func_name} in {last_filename}:"
f"{last_lineno}\n")
else:
f.write(f"{ts} Return from"
f" {func_name} in {filename}:{lineno}"
f" to {last_func_name} in {last_filename}:"
f"{last_lineno}\n")
except NameError:
# modules are deleted during shutdown
pass
return partial(_trace_calls, log_path, root_dir)
def enable_trace_function_call(log_file_path: str,
root_dir: Optional[str] = None):
"""
Enable tracing of every function call in code under `root_dir`.
This is useful for debugging hangs or crashes.
`log_file_path` is the path to the log file.
`root_dir` is the root directory of the code to trace. If None, it is the
fastvideo root directory.
Note that this call is thread-level, any threads calling this function
will have the trace enabled. Other threads will not be affected.
"""
logger.warning(
"FASTVIDEO_TRACE_FUNCTION is enabled. It will record every"
" function executed by Python. This will slow down the code. It "
"is suggested to be used for debugging hang or crashes only.")
logger.info("Trace frame log is saved to %s", log_file_path)
if root_dir is None:
# by default, this is the fastvideo root directory
root_dir = os.path.dirname(os.path.dirname(__file__))
sys.settrace(partial(_trace_calls, log_file_path, root_dir))
-8
View File
@@ -1,8 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.v1.logging_utils.formatter import NewLineFormatter
__all__ = [
"NewLineFormatter",
]
-18
View File
@@ -1,18 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from vllm
import logging
class NewLineFormatter(logging.Formatter):
"""Adds logging prefix to newlines to align multi-line messages."""
def __init__(self, fmt, datefmt=None, style="%"):
logging.Formatter.__init__(self, fmt, datefmt, style)
def format(self, record):
msg = logging.Formatter.format(self, record)
if record.message != "":
parts = msg.split(record.message)
msg = msg.replace("\n", "\r\n" + parts[0])
return msg
-31
View File
@@ -1,31 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import torch.nn as nn
from typing import Dict
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def get_scheduler(module_path: str, architecture: str, inference_args: InferenceArgs) -> Dict:
"""Create a scheduler based on the inference args. Can be overridden by subclasses."""
if hasattr(inference_args, 'denoise_type') and inference_args.denoise_type == "flow":
# TODO(will): add schedulers to register or create a new scheduler registry
# TODO(will): default to config file but allow override through
# inference args. Currently only uses inference args.
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchDiscreteScheduler
return FlowMatchDiscreteScheduler(
shift=inference_args.flow_shift,
solver=inference_args.flow_solver,
)
else:
raise ValueError(f"Invalid denoise type: {inference_args.denoise_type}")
__all__ = [
"set_random_seed",
"BasevLLMParameter",
"PackedvLLMParameter",
"get_model",
"get_scheduler",
]
-12
View File
@@ -1,12 +0,0 @@
from torch import nn
class BaseDiT(nn.Module):
_fsdp_shard_conditions = []
attention_head_dim: int = None
def __init__(self, *args, **kwargs):
super().__init__()
def forward(self, *args, **kwargs):
pass
-897
View File
@@ -1,897 +0,0 @@
import torch
import torch.nn as nn
from typing import Optional, Tuple, List
from fastvideo.v1.attention.flash_attn import DistributedAttention, LocalAttention
from fastvideo.v1.layers.linear import ReplicatedLinear
from fastvideo.v1.layers.layernorm import LayerNormScaleShift, ScaleResidual, ScaleResidualLayerNormScaleShift
from fastvideo.v1.layers.visual_embedding import PatchEmbed, TimestepEmbedder, ModulateProjection, unpatchify
from fastvideo.v1.layers.rotary_embedding import _apply_rotary_emb, get_rotary_pos_embed
from fastvideo.v1.distributed.parallel_state import get_sequence_model_parallel_world_size, get_sequence_model_parallel_rank
# TODO(will-PY-refactor): RMSNorm ....
from fastvideo.v1.layers.mlp import MLP
from fastvideo.v1.models.dits.base import BaseDiT
from diffusers.loaders import PeftAdapterMixin
class HunyuanRMSNorm(nn.Module):
def __init__(
self,
dim: int,
elementwise_affine=True,
eps: float = 1e-6,
device=None,
dtype=None,
):
"""
Initialize the RMSNorm normalization layer.
Args:
dim (int): The dimension of the input tensor.
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
Attributes:
eps (float): A small value added to the denominator for numerical stability.
weight (nn.Parameter): Learnable scaling parameter.
"""
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
def _norm(self, x):
"""
Apply the RMSNorm normalization to the input tensor.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The normalized tensor.
"""
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
"""
Forward pass through the RMSNorm layer.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The output tensor after applying RMSNorm.
"""
output = self._norm(x.float()).type_as(x)
if hasattr(self, "weight"):
output = output * self.weight
return output
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal DiT block with separate modulation for text and image/video,
using distributed attention and linear layers.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float,
dtype: Optional[torch.dtype] = None,
):
super().__init__()
self.deterministic = False
self.num_attention_heads = num_attention_heads
head_dim = hidden_size // num_attention_heads
mlp_hidden_dim = int(hidden_size * mlp_ratio)
# Image modulation components
self.img_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
)
# Fused operations for image stream
self.img_attn_norm = LayerNormScaleShift(hidden_size, norm_type="layer", elementwise_affine=False, dtype=dtype)
self.img_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(hidden_size, norm_type="layer", elementwise_affine=False, dtype=dtype)
self.img_mlp_residual = ScaleResidual()
# Image attention components
# self.img_attn_qkv = ReplicatedLinear(
# hidden_size,
# hidden_size * 3,
# bias=True,
# params_dtype=dtype
# )
for name in ["q", "k", "v"]:
setattr(self, f"img_attn_to_{name}", ReplicatedLinear(
hidden_size, hidden_size, bias=True, params_dtype=dtype
))
self.img_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.img_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.img_attn_proj = ReplicatedLinear(
hidden_size,
hidden_size,
bias=True,
params_dtype=dtype
)
self.img_mlp = MLP(
hidden_size,
mlp_hidden_dim,
bias=True,
dtype=dtype
)
# Text modulation components
self.txt_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
)
# Fused operations for text stream
self.txt_attn_norm = LayerNormScaleShift(hidden_size, norm_type="layer", elementwise_affine=False, dtype=dtype)
self.txt_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(hidden_size, norm_type="layer", elementwise_affine=False, dtype=dtype)
self.txt_mlp_residual = ScaleResidual()
# Text attention components
self.txt_attn_qkv = ReplicatedLinear(
hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype
)
# QK norm layers for text
self.txt_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_proj = ReplicatedLinear(
hidden_size,
hidden_size,
bias=True,
params_dtype=dtype
)
self.txt_mlp = MLP(
hidden_size,
mlp_hidden_dim,
bias=True,
dtype=dtype
)
# Distributed attention
self.attn = DistributedAttention(
dropout_rate=0.0,
causal=False
)
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
vec: torch.Tensor,
freqs_cis: tuple = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
# Process modulation vectors
img_mod_outputs = self.img_mod(vec)
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=-1)
txt_mod_outputs = self.txt_mod(vec)
(
txt_attn_shift,
txt_attn_scale,
txt_attn_gate,
txt_mlp_shift,
txt_mlp_scale,
txt_mlp_gate,
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
# Prepare image for attention using fused operation
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale)
# Get QKV for image
# img_qkv, _ = self.img_attn_qkv(img_attn_input)
img_q_, _ = self.img_attn_to_q(img_attn_input)
img_k_, _ = self.img_attn_to_k(img_attn_input)
img_v_, _ = self.img_attn_to_v(img_attn_input)
img_qkv = torch.cat((img_q_, img_k_, img_v_), dim=-1)
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
# batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
# Split QKV
img_qkv = img_qkv.view(batch_size, image_seq_len, 3, self.num_attention_heads, -1)
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :, 2]
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply rotary embeddings
cos, sin = freqs_cis
img_q, img_k = _apply_rotary_emb(img_q, cos, sin, is_neox_style=False), _apply_rotary_emb(img_k, cos, sin, is_neox_style=False)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
# Get QKV for text
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
# Split QKV
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3, self.num_attention_heads, -1)
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :, 2]
# Apply QK-Norm if needed
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
# Run distributed attention
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v)
img_attn_out, _ = self.img_attn_proj(img_attn.view(batch_size, image_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale
)
# Process image MLP
img_mlp_out = self.img_mlp(img_mlp_input)
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
# Process text attention output
txt_attn_out, _ = self.txt_attn_proj(txt_attn.reshape(batch_size, text_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale
)
# Process text MLP
txt_mlp_out = self.txt_mlp(txt_mlp_input)
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
return img, txt
class MMSingleStreamBlock(nn.Module):
"""
A DiT block with parallel linear layers using distributed attention
and tensor parallelism.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
):
super().__init__()
self.deterministic = False
self.hidden_size = hidden_size
self.num_attention_heads = num_attention_heads
head_dim = hidden_size // num_attention_heads
mlp_hidden_dim = int(hidden_size * mlp_ratio)
self.mlp_hidden_dim = mlp_hidden_dim
# Combined QKV and MLP input projection
# self.linear1 = ReplicatedLinear(
# hidden_size,
# hidden_size * 3 + mlp_hidden_dim,
# bias=True,
# params_dtype=dtype
# )
# self.linear1_qkv = ReplicatedLinear(
# hidden_size,
# hidden_size * 3,
# bias=True,
# params_dtype=dtype
# )
for name in ["q", "k", "v"]:
setattr(self, f"linear1_to_{name}", ReplicatedLinear(
hidden_size, hidden_size, bias=True, params_dtype=dtype
))
self.linear1_mlp = ReplicatedLinear(
hidden_size,
mlp_hidden_dim,
bias=True,
params_dtype=dtype
)
# Combined projection and MLP output
self.linear2 = ReplicatedLinear(
hidden_size + mlp_hidden_dim,
hidden_size,
bias=True,
params_dtype=dtype
)
# QK norm layers
self.q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
# Fused operations with better naming
self.input_norm_scale_shift = LayerNormScaleShift(hidden_size, norm_type="layer", eps=1e-6, elementwise_affine=False, dtype=dtype)
self.output_residual = ScaleResidual()
# Activation function
self.mlp_act = nn.GELU(approximate="tanh")
# Modulation
self.modulation = ModulateProjection(
hidden_size,
factor=3,
act_layer="silu",
dtype=dtype
)
# Distributed attention
self.attn = DistributedAttention(
dropout_rate=0.0,
causal=False
)
def forward(
self,
x: torch.Tensor,
vec: torch.Tensor,
txt_len: int,
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
) -> torch.Tensor:
# Process modulation
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
# Apply pre-norm and modulation using fused operation
x_mod = self.input_norm_scale_shift(x, mod_shift, mod_scale)
# Get combined projections
linear1_q, _ = self.linear1_to_q(x_mod)
linear1_k, _ = self.linear1_to_k(x_mod)
linear1_v, _ = self.linear1_to_v(x_mod)
linear1_qkv = torch.cat((linear1_q, linear1_k, linear1_v), dim=-1)
linear1_mlp, _ = self.linear1_mlp(x_mod)
linear1_out = torch.cat((linear1_qkv, linear1_mlp), dim=-1)
# Split into QKV and MLP parts
qkv, mlp = torch.split(
linear1_out, [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
)
# Process QKV
batch_size, seq_len = qkv.shape[0], qkv.shape[1]
qkv = qkv.view(batch_size, seq_len, 3, self.num_attention_heads, -1)
q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2]
# Apply QK-Norm
q = self.q_norm(q).to(v.dtype)
k = self.k_norm(k).to(v.dtype)
# Split into image and text parts
img_q, txt_q = q[:, :-txt_len], q[:, -txt_len:]
img_k, txt_k = k[:, :-txt_len], k[:, -txt_len:]
img_v, txt_v = v[:, :-txt_len], v[:, -txt_len:]
# Apply rotary embeddings to image parts
cos, sin = freqs_cis
img_q, img_k = _apply_rotary_emb(img_q, cos, sin, is_neox_style=False), _apply_rotary_emb(img_k, cos, sin, is_neox_style=False)
# Run distributed attention
img_attn_output, txt_attn_output = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v)
attn_output = torch.cat((img_attn_output, txt_attn_output), dim=1).view(batch_size, seq_len, -1)
# Process MLP activation
mlp_output = self.mlp_act(mlp)
# Combine attention and MLP outputs
combined = torch.cat((attn_output, mlp_output), dim=-1)
# Final projection
output, _ = self.linear2(combined)
# Apply residual connection with gating using fused operation
return self.output_residual(x, output, mod_gate)
class HunyuanVideoTransformer3DModel(BaseDiT, PeftAdapterMixin):
"""
HunyuanVideo Transformer backbone adapted for distributed training.
This implementation uses distributed attention and linear layers for efficient
parallel processing across multiple GPUs.
Based on the architecture from:
- Flux.1: https://github.com/black-forest-labs/flux
- MMDiT: http://arxiv.org/abs/2403.03206
"""
# PY: we make the input args the same as HF config
# shard single stream, double stream blocks, and refiner_blocks
_fsdp_shard_conditions = [
lambda n, m: "double" in n and str.isdigit(n.split(".")[-1]),
lambda n, m: "single" in n and str.isdigit(n.split(".")[-1]),
lambda n, m: "refiner" in n and str.isdigit(n.split(".")[-1]),
]
_param_names_mapping = {
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$": r"txt_in.t_embedder.mlp.fc_in.\1",
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$": r"txt_in.t_embedder.mlp.fc_out.\1",
r"^context_embedder\.proj_in\.(.*)$": r"txt_in.input_embedder.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$": r"txt_in.c_embedder.fc_in.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$": r"txt_in.c_embedder.fc_out.\1",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$": r"txt_in.refiner_blocks.\1.norm1.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$": r"txt_in.refiner_blocks.\1.norm2.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$": r"txt_in.refiner_blocks.\1.self_attn_to_q.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$": r"txt_in.refiner_blocks.\1.self_attn_to_k.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$": r"txt_in.refiner_blocks.\1.self_attn_to_v.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$": r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$": r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$": r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$": r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
# 3. x_embedder mapping:
r"^x_embedder\.proj\.(.*)$": r"img_in.proj.\1",
# 4. Top-level time_text_embed mappings:
r"^time_text_embed\.timestep_embedder\.linear_1\.(.*)$": r"time_in.mlp.fc_in.\1",
r"^time_text_embed\.timestep_embedder\.linear_2\.(.*)$": r"time_in.mlp.fc_out.\1",
r"^time_text_embed\.guidance_embedder\.linear_1\.(.*)$": r"guidance_in.mlp.fc_in.\1",
r"^time_text_embed\.guidance_embedder\.linear_2\.(.*)$": r"guidance_in.mlp.fc_out.\1",
r"^time_text_embed\.text_embedder\.linear_1\.(.*)$": r"vector_in.fc_in.\1",
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$": r"vector_in.fc_out.\1",
# 5. transformer_blocks mapping:
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$": r"double_blocks.\1.img_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$": r"double_blocks.\1.txt_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$": r"double_blocks.\1.img_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$": r"double_blocks.\1.img_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": r"double_blocks.\1.img_attn_to_q.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$": r"double_blocks.\1.img_attn_to_k.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$": r"double_blocks.\1.img_attn_to_v.\2",
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$": r"double_blocks.\1.img_attn_proj.\2",
# Corrected: merge attn.to_add_out into the main projection.
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$": r"double_blocks.\1.txt_attn_proj.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$": r"double_blocks.\1.txt_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$": r"double_blocks.\1.txt_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$": r"double_blocks.\1.img_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$": r"double_blocks.\1.img_mlp.fc_out.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$": r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$": r"double_blocks.\1.txt_mlp.fc_out.\2",
# 6. single_transformer_blocks mapping:
r"^single_transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$": r"single_blocks.\1.q_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$": r"single_blocks.\1.k_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": r"single_blocks.\1.linear1_to_q.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$": r"single_blocks.\1.linear1_to_k.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$": r"single_blocks.\1.linear1_to_v.\2",
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$": r"single_blocks.\1.linear1_mlp.\2",
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$": r"single_blocks.\1.linear2.\2",
r"^single_transformer_blocks\.(\d+)\.norm\.linear\.(.*)$": r"single_blocks.\1.modulation.linear.\2",
# 7. Final layers mapping:
r"^norm_out\.linear\.(.*)$": r"final_layer.adaLN_modulation.linear.\1",
r"^proj_out\.(.*)$": r"final_layer.linear.\1",
}
def __init__(
self,
patch_size: int = 2,
patch_size_t: int = 1,
in_channels: int = 16,
out_channels: int = 16,
num_attention_heads: int = 24,
attention_head_dim: int = 128,
mlp_ratio: float = 4.0,
num_layers: int = 20,
num_single_layers: int = 40,
num_refiner_layers: int = 2,
rope_axes_dim: List[int] = [16, 56, 56],
guidance_embeds: bool = False,
dtype: Optional[torch.dtype] = None,
text_embed_dim: int = 4096,
pooled_projection_dim: int = 768,
rope_theta: int = 256,
qk_norm: str = "rms_norm", #TODO(PY)
):
super().__init__()
hidden_size = attention_head_dim * num_attention_heads
self.patch_size = [patch_size_t, patch_size, patch_size]
self.in_channels = in_channels
self.out_channels = in_channels if out_channels is None else out_channels
self.unpatchify_channels = self.out_channels
self.guidance_embeds = guidance_embeds
self.rope_dim_list = rope_axes_dim
self.rope_theta = rope_theta
self.text_states_dim = text_embed_dim
self.text_states_dim_2 = pooled_projection_dim
if hidden_size % num_attention_heads != 0:
raise ValueError(f"Hidden size {hidden_size} must be divisible by num_attention_heads {num_attention_heads}")
pe_dim = hidden_size // num_attention_heads
if sum(rope_axes_dim) != pe_dim:
raise ValueError(f"Got {rope_axes_dim} but expected positional dim {pe_dim}")
self.hidden_size = hidden_size
self.num_attention_heads = num_attention_heads
# Image projection
self.img_in = PatchEmbed(
self.patch_size,
self.in_channels,
self.hidden_size,
dtype=dtype
)
self.txt_in = SingleTokenRefiner(
self.text_states_dim,
hidden_size,
num_attention_heads,
depth=num_refiner_layers,
dtype=dtype
)
# Time modulation
self.time_in = TimestepEmbedder(
self.hidden_size,
act_layer="silu",
dtype=dtype
)
# Text modulation
self.vector_in = MLP(
self.text_states_dim_2,
self.hidden_size,
self.hidden_size,
act_type="silu",
dtype=dtype
)
# Guidance modulation
self.guidance_in = (
TimestepEmbedder(self.hidden_size, act_layer="silu", dtype=dtype)
if guidance_embeds else None
)
# Double blocks
self.double_blocks = nn.ModuleList([
MMDoubleStreamBlock(
hidden_size,
num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
) for _ in range(num_layers)
])
# Single blocks
self.single_blocks = nn.ModuleList([
MMSingleStreamBlock(
hidden_size,
num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
) for _ in range(num_single_layers)
])
self.final_layer = FinalLayer(
hidden_size,
self.patch_size,
self.out_channels,
dtype=dtype
)
# TODO: change the input the FORWAD_BACTCH Dict
# TODO: change output to a dict
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor,
guidance=None,
):
"""
Forward pass of the HunyuanDiT model.
Args:
hidden_states: Input image/video latents [B, C, T, H, W]
encoder_hidden_states: Text embeddings [B, L, D]
timestep: Diffusion timestep
guidance: Guidance scale for CFG
Returns:
Tuple of (output)
"""
if guidance is None:
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=hidden_states.dtype)
img = x = hidden_states
t = timestep
# Split text embeddings - first token is global, rest are per-token
txt = encoder_hidden_states[:, 1:]
text_states_2 = encoder_hidden_states[:, 0, :self.text_states_dim_2]
# Get spatial dimensions
_, _, ot, oh, ow = x.shape
tt, th, tw = (
ot // self.patch_size[0],
oh // self.patch_size[1],
ow // self.patch_size[2],
)
# Get rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed((tt * get_sequence_model_parallel_world_size(), th, tw), self.hidden_size, self.num_attention_heads, self.rope_dim_list, self.rope_theta)
freqs_cos = freqs_cos.to(x.device)
freqs_sin = freqs_sin.to(x.device)
# Prepare modulation vectors
vec = self.time_in(t)
# Add text modulation
vec = vec + self.vector_in(text_states_2)
# Add guidance modulation if needed
if self.guidance_embeds and guidance is not None:
vec = vec + self.guidance_in(guidance)
# Embed image and text
img = self.img_in(img)
txt = self.txt_in(txt, t)
txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1]
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# Process through double stream blocks
for index, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis]
img, txt = block(*double_block_args)
# Merge txt and img to pass through single stream blocks
x = torch.cat((img, txt), 1)
# Process through single stream blocks
if len(self.single_blocks) > 0:
for index, block in enumerate(self.single_blocks):
single_block_args = [
x,
vec,
txt_seq_len,
freqs_cis,
]
x = block(*single_block_args)
# Extract image features
img = x[:, :img_seq_len, ...]
# Final layer processing
img = self.final_layer(img, vec)
# Unpatchify to get original shape
img = unpatchify(img, tt, th, tw, self.patch_size, self.out_channels)
return img
class SingleTokenRefiner(nn.Module):
"""
A token refiner that processes text embeddings with attention to improve
their representation for cross-attention with image features.
"""
def __init__(
self,
in_channels,
hidden_size,
num_attention_heads,
depth=2,
qkv_bias=True,
dtype=None,
):
super().__init__()
# Input projection
self.input_embedder = ReplicatedLinear(
in_channels,
hidden_size,
bias=True,
params_dtype=dtype
)
# Timestep embedding
self.t_embedder = TimestepEmbedder(
hidden_size,
act_layer="silu",
dtype=dtype
)
# Context embedding
self.c_embedder = MLP(
in_channels,
hidden_size,
hidden_size,
act_type="silu",
dtype=dtype
)
# Refiner blocks
self.refiner_blocks = nn.ModuleList([
IndividualTokenRefinerBlock(
hidden_size,
num_attention_heads,
qkv_bias=qkv_bias,
dtype=dtype,
) for _ in range(depth)
])
def forward(self, x, t):
# Get timestep embeddings
timestep_aware_representations = self.t_embedder(t)
# Get context-aware representations
context_aware_representations = torch.mean(x, dim=1)
context_aware_representations = self.c_embedder(context_aware_representations)
c = timestep_aware_representations + context_aware_representations
# Project input
x, _ = self.input_embedder(x)
# Process through refiner blocks
for block in self.refiner_blocks:
x = block(x, c)
return x
class IndividualTokenRefinerBlock(nn.Module):
"""
A transformer block for refining individual tokens with self-attention.
"""
def __init__(
self,
hidden_size,
num_attention_heads,
mlp_ratio=4.0,
qkv_bias=True,
dtype=None,
):
super().__init__()
self.num_attention_heads = num_attention_heads
mlp_hidden_dim = int(hidden_size * mlp_ratio)
# Normalization and attention
self.norm1 = nn.LayerNorm(hidden_size, eps=1e-6, elementwise_affine=True, dtype=dtype)
# self.self_attn_qkv = ReplicatedLinear(
# hidden_size,
# hidden_size * 3,
# bias=qkv_bias,
# params_dtype=dtype
# )
for name in ["q", "k", "v"]:
setattr(self, f"self_attn_to_{name}", ReplicatedLinear(
hidden_size, hidden_size, bias=True, params_dtype=dtype
))
self.self_attn_proj = ReplicatedLinear(
hidden_size,
hidden_size,
bias=qkv_bias,
params_dtype=dtype
)
# MLP
self.norm2 = nn.LayerNorm(hidden_size, eps=1e-6, elementwise_affine=True, dtype=dtype)
self.mlp = MLP(
hidden_size,
mlp_hidden_dim,
bias=True,
act_type="silu",
dtype=dtype
)
# Modulation
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype
)
# Scaled dot product attention
self.attn = LocalAttention()
def forward(self, x, c):
# Get modulation parameters
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=-1)
# Self-attention
norm_x = self.norm1(x)
q_, _ = self.self_attn_to_q(norm_x)
k_, _ = self.self_attn_to_k(norm_x)
v_, _ = self.self_attn_to_v(norm_x)
qkv = torch.cat((q_, k_, v_), dim=-1)
batch_size, seq_len = qkv.shape[0], qkv.shape[1]
qkv = qkv.view(batch_size, seq_len, 3, self.num_attention_heads, -1)
q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2]
# Run scaled dot product attention
attn_output = self.attn(q, k, v) # [B, L, H, D]
attn_output = attn_output.reshape(batch_size, seq_len, -1) # [B, L, H*D]
# Project and apply residual connection with gating
attn_out, _ = self.self_attn_proj(attn_output)
x = x + attn_out * gate_msa.unsqueeze(1)
# MLP
mlp_out = self.mlp(self.norm2(x))
x = x + mlp_out * gate_mlp.unsqueeze(1)
return x
class FinalLayer(nn.Module):
"""
The final layer of DiT that projects features to pixel space.
"""
def __init__(self, hidden_size, patch_size, out_channels, dtype=None):
super().__init__()
# Normalization
self.norm_final = nn.LayerNorm(hidden_size, eps=1e-6, elementwise_affine=False, dtype=dtype)
output_dim = patch_size[0] * patch_size[1] * patch_size[2] * out_channels
self.linear = ReplicatedLinear(
hidden_size,
output_dim,
bias=True,
params_dtype=dtype
)
# Modulation
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype
)
def forward(self, x, c):
# What the heck HF? Why you change the scale and shift order here???
scale, shift = self.adaLN_modulation(c).chunk(2, dim=-1)
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
x, _ = self.linear(x)
return x
-456
View File
@@ -1,456 +0,0 @@
import torch
import torch.nn as nn
import math
from typing import Optional, Tuple, List, Union, Dict, Any
from fastvideo.v1.attention.flash_attn import DistributedAttention, LocalAttention
from fastvideo.v1.layers.linear import ReplicatedLinear
from fastvideo.v1.layers.layernorm import LayerNormScaleShift, ScaleResidual, ScaleResidualLayerNormScaleShift, RMSNorm
from fastvideo.v1.layers.visual_embedding import PatchEmbed, TimestepEmbedder, ModulateProjection
from fastvideo.v1.layers.rotary_embedding import _apply_rotary_emb, get_rotary_pos_embed
from fastvideo.v1.distributed.parallel_state import get_sequence_model_parallel_world_size
# from torch.nn import RMSNorm
# TODO: RMSNorm ....
from fastvideo.v1.layers.mlp import MLP
from fastvideo.v1.models.dits.base import BaseDiT
class WanImageEmbedding(torch.nn.Module):
def __init__(self, in_features: int, out_features: int):
super().__init__()
self.norm1 = nn.LayerNorm(in_features)
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
self.norm2 = nn.LayerNorm(out_features)
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
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,
text_embed_dim: int,
image_embed_dim: Optional[int] = None,
):
super().__init__()
self.time_embedder = TimestepEmbedder(dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
self.time_modulation = ModulateProjection(dim, factor=6, act_layer="silu")
self.text_embedder = MLP(text_embed_dim, dim, dim, bias=True, act_type="gelu_pytorch_tanh")
self.image_embedder = None
if image_embed_dim is not None:
self.image_embedder = WanImageEmbedding(image_embed_dim, dim)
def forward(
self,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
):
with torch.cuda.amp.autocast(dtype=torch.float32):
temb = self.time_embedder(timestep.float())
timestep_proj = self.time_modulation(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 WanSelfAttention(nn.Module):
def __init__(self,
dim,
num_heads,
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
parallel_attention=False):
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.window_size = window_size
self.qk_norm = qk_norm
self.eps = eps
self.parallel_attention = parallel_attention
# layers
self.to_q = ReplicatedLinear(dim, dim)
self.to_k = ReplicatedLinear(dim, dim)
self.to_v = ReplicatedLinear(dim, dim)
self.to_out = ReplicatedLinear(dim, dim)
self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
# Scaled dot product attention
self.attn = LocalAttention(dropout_rate=0, softmax_scale=None, causal=False)
def forward(self, x, seq_lens, grid_sizes, freqs):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
seq_lens(Tensor): Shape [B]
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
"""
pass
class WanT2VCrossAttention(WanSelfAttention):
def forward(self, x, context, context_lens):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
context_lens(Tensor): Shape [B]
"""
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
x = self.attn(q, k, v)
# output
x = x.flatten(2)
x, _ = self.to_out(x)
return x
class WanI2VCrossAttention(WanSelfAttention):
def __init__(self,
dim,
num_heads,
window_size=(-1, -1),
qk_norm=True,
eps=1e-6):
super().__init__(dim, num_heads, window_size, qk_norm, eps)
self.add_k_proj = ReplicatedLinear(dim, dim)
self.add_v_proj = ReplicatedLinear(dim, dim)
self.norm_added_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_added_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
def forward(self, x, context, context_lens):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
context_lens(Tensor): Shape [B]
"""
context_img = context[:, :257]
context = context[:, 257:]
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
k_img = self.norm_added_k.forward_native(self.add_k_proj(context_img)[0]).view(b, -1, n, d)
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
img_x = self.attn(q, k_img, v_img)
# compute attention
x = self.attn(q, k, v)
# output
x = x.flatten(2)
img_x = img_x.flatten(2)
x = x + img_x
x, _ = self.to_out(x)
return x
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 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = DistributedAttention(
dropout_rate=0.0,
causal=False
)
self.hidden_dim = dim
self.num_attention_heads = num_heads
dim_head = dim // num_heads
if qk_norm == "rms_norm":
self.norm_q = RMSNorm(dim_head, eps=eps)
self.norm_k = RMSNorm(dim_head, eps=eps)
elif qk_norm == "rms_norm_across_heads":
# LTX applies qk norm across all heads
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
else:
print("QK Norm type not supported")
raise Exception
assert cross_attn_norm is True
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(dim, norm_type="layer", eps=eps, elementwise_affine=True, dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
# I2V
self.attn2 = WanI2VCrossAttention(dim, num_heads, qk_norm=qk_norm, eps=eps)
else:
# T2V
self.attn2 = WanT2VCrossAttention(dim, num_heads, qk_norm=qk_norm, eps=eps)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(dim, norm_type="layer", eps=eps, elementwise_affine=False, dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
self.mlp_residual = ScaleResidual()
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,
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
assert temb.dtype == torch.float32
with torch.cuda.amp.autocast(dtype=torch.float32):
e = self.scale_shift_table + temb
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = self.norm1(hidden_states.float()).to(dtype=orig_dtype) * (1 + scale_msa) + shift_msa
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
# Apply rotary embeddings
cos, sin = freqs_cis
query, key = _apply_rotary_emb(query, cos, sin, is_neox_style=False), _apply_rotary_emb(key, cos, sin, is_neox_style=False)
attn_output, _ = self.attn1(query, key, value)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(hidden_states, attn_output, gate_msa, null_shift, null_scale)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states, context=encoder_hidden_states, context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
return hidden_states
class WanTransformer3DModel(BaseDiT):
_fsdp_shard_conditions = [
lambda n, m: "blocks" in n and str.isdigit(n.split(".")[-1]),
]
_param_names_mapping = {
r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1",
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$": r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$": r"condition_embedder.image_embedder.ff.fc_in.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$": r"condition_embedder.image_embedder.ff.fc_out.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$": r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$": r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$": r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$": r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
r"blocks\.(\d+)\.norm2\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
}
def __init__(
self,
patch_size: Tuple[int] = (1, 2, 2),
text_len = 512,
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,
) -> None:
super().__init__()
inner_dim = num_attention_heads * attention_head_dim
self.inner_dim = inner_dim
self.num_attention_heads = num_attention_heads
self.out_channels = out_channels or in_channels
self.patch_size = patch_size
self.text_len = text_len
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=in_channels, embed_dim=inner_dim, patch_size=patch_size, flatten=False)
# 2. Condition embeddings
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=freq_dim,
text_embed_dim=text_dim,
image_embed_dim=image_dim,
)
# 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 = LayerNormScaleShift(inner_dim, norm_type="layer", eps=eps, elementwise_affine=False, dtype=torch.float32)
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,
seq_len: Optional[int] = None,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
y: Optional[torch.Tensor] = None,
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if y is not None:
hidden_states = torch.cat([hidden_states, y], dim=1)
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# Get rotary embeddings
d = self.inner_dim // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed((post_patch_num_frames * get_sequence_model_parallel_world_size(), post_patch_height, post_patch_width), self.inner_dim, self.num_attention_heads, rope_dim_list, rope_theta=10000)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
if seq_len is None:
seq_len = hidden_states.size(1)
hidden_states = torch.cat([hidden_states, hidden_states.new_zeros(1, seq_len - hidden_states.size(1), hidden_states.size(2))], dim=1)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
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, freqs_cis
)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, freqs_cis)
# 5. Output norm, projection & unpatchify
with torch.cuda.amp.autocast(dtype=torch.float32):
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
hidden_states = self.proj_out(hidden_states)
output = self.unpatchify(hidden_states, grid_sizes)
return output.float()
def unpatchify(self, x, grid_sizes):
r"""
Reconstruct video tensors from patch embeddings.
Args:
x (List[Tensor]):
List of patchified features, each with shape [L, C_out * prod(patch_size)]
grid_sizes (Tensor):
Original spatial-temporal grid dimensions before patching,
shape [B, 3] (3 dimensions correspond to F_patches, H_patches, W_patches)
Returns:
Tensor:
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
"""
c = self.out_channels
out = []
for u, v in zip(x, grid_sizes.tolist()):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = u.permute(6, 0, 3, 1, 4, 2, 5)
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
out.append(u)
out = torch.cat(out, dim=0)
return out
-656
View File
@@ -1,656 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Minimal implementation of CLIPVisionModel intended to be only used
within a vision language model."""
from typing import Iterable, Optional, Set, Tuple, Union
import torch
import torch.nn as nn
from transformers import CLIPVisionConfig, CLIPTextConfig
from transformers.modeling_outputs import BaseModelOutputWithPooling
# from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask
from vllm.attention.layer import MultiHeadAttention
# from fastvideo.v1.attention.flash_attn import LocalAttention
from fastvideo.v1.distributed import divide, get_tensor_model_parallel_world_size
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.layers.linear import (ColumnParallelLinear,
QKVParallelLinear,
RowParallelLinear)
# TODO: support quantization
# from vllm.model_executor.layers.quantization import QuantizationConfig
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
from vllm.model_executor.models.interfaces import SupportsQuant
from .vision import VisionEncoderInfo, resolve_visual_encoder_outputs
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class QuantizationConfig:
pass
class CLIPEncoderInfo(VisionEncoderInfo[CLIPVisionConfig]):
def get_num_image_tokens(
self,
*,
image_width: int,
image_height: int,
) -> int:
return self.get_patch_grid_length()**2 + 1
def get_max_image_tokens(self) -> int:
return self.get_patch_grid_length()**2 + 1
def get_image_size(self) -> int:
return self.vision_config.image_size
def get_patch_size(self) -> int:
return self.vision_config.patch_size
def get_patch_grid_length(self) -> int:
image_size, patch_size = self.get_image_size(), self.get_patch_size()
assert image_size % patch_size == 0
return image_size // patch_size
# Adapted from https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py#L164 # noqa
class CLIPVisionEmbeddings(nn.Module):
def __init__(self, config: CLIPVisionConfig):
super().__init__()
self.config = config
self.embed_dim = config.hidden_size
self.image_size = config.image_size
self.patch_size = config.patch_size
assert self.image_size % self.patch_size == 0
self.class_embedding = nn.Parameter(torch.randn(self.embed_dim))
self.patch_embedding = nn.Conv2d(
in_channels=config.num_channels,
out_channels=self.embed_dim,
kernel_size=self.patch_size,
stride=self.patch_size,
bias=False,
)
self.num_patches = (self.image_size // self.patch_size)**2
self.num_positions = self.num_patches + 1
self.position_embedding = nn.Embedding(self.num_positions,
self.embed_dim)
self.register_buffer("position_ids",
torch.arange(self.num_positions).expand((1, -1)),
persistent=False)
def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:
batch_size = pixel_values.shape[0]
target_dtype = self.patch_embedding.weight.dtype
patch_embeds = self.patch_embedding(pixel_values.to(
dtype=target_dtype)) # shape = [*, width, grid, grid]
patch_embeds = patch_embeds.flatten(2).transpose(1, 2)
class_embeds = self.class_embedding.expand(batch_size, 1, -1)
embeddings = torch.cat([class_embeds, patch_embeds], dim=1)
embeddings = embeddings + self.position_embedding(self.position_ids)
return embeddings
class CLIPTextEmbeddings(nn.Module):
def __init__(self, config: CLIPTextConfig):
super().__init__()
self.config = config
embed_dim = config.hidden_size
self.token_embedding = nn.Embedding(config.vocab_size, embed_dim)
self.position_embedding = nn.Embedding(config.max_position_embeddings, embed_dim)
# position_ids (1, len position emb) is contiguous in memory and exported when serialized
self.register_buffer(
"position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)), persistent=False
)
def forward(
self,
input_ids: Optional[torch.LongTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
) -> torch.Tensor:
seq_length = input_ids.shape[-1] if input_ids is not None else inputs_embeds.shape[-2]
max_position_embedding = self.position_embedding.weight.shape[0]
if seq_length > max_position_embedding:
raise ValueError(
f"Sequence length must be less than max_position_embeddings (got `sequence length`: "
f"{seq_length} and max_position_embeddings: {max_position_embedding}"
)
if position_ids is None:
position_ids = self.position_ids[:, :seq_length]
if inputs_embeds is None:
inputs_embeds = self.token_embedding(input_ids)
position_embeddings = self.position_embedding(position_ids)
embeddings = inputs_embeds + position_embeddings
return embeddings
class CLIPAttention(nn.Module):
"""Multi-headed attention from 'Attention Is All You Need' paper"""
def __init__(
self,
config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
self.config = config
self.embed_dim = config.hidden_size
self.num_heads = config.num_attention_heads
self.head_dim = self.embed_dim // self.num_heads
if self.head_dim * self.num_heads != self.embed_dim:
raise ValueError(
"embed_dim must be divisible by num_heads "
f"(got `embed_dim`: {self.embed_dim} and `num_heads`:"
f" {self.num_heads}).")
self.scale = self.head_dim**-0.5
self.dropout = config.attention_dropout
self.qkv_proj = QKVParallelLinear(
hidden_size=self.embed_dim,
head_size=self.head_dim,
total_num_heads=self.num_heads,
quant_config=quant_config,
prefix=f"{prefix}.qkv_proj",
)
self.out_proj = RowParallelLinear(
input_size=self.embed_dim,
output_size=self.embed_dim,
quant_config=quant_config,
prefix=f"{prefix}.out_proj",
)
self.tp_size = get_tensor_model_parallel_world_size()
self.num_heads_per_partition = divide(self.num_heads, self.tp_size)
self.attn = MultiHeadAttention(self.num_heads_per_partition,
self.head_dim, self.scale)
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
return tensor.view(bsz, seq_len, self.num_heads,
self.head_dim).transpose(1, 2).contiguous()
def forward(
self,
hidden_states: torch.Tensor,
):
"""Input shape: Batch x Time x Channel"""
qkv_states, _ = self.qkv_proj(hidden_states)
query_states, key_states, value_states = qkv_states.chunk(3, dim=-1)
# use flash_attn_func
from flash_attn import flash_attn_func
query_states = query_states.reshape(query_states.shape[0], query_states.shape[1], self.num_heads_per_partition, self.head_dim)
key_states = key_states.reshape(key_states.shape[0], key_states.shape[1], self.num_heads_per_partition, self.head_dim)
value_states = value_states.reshape(value_states.shape[0], value_states.shape[1], self.num_heads_per_partition, self.head_dim)
attn_output = flash_attn_func(query_states, key_states, value_states,softmax_scale=self.scale, causal=True)
attn_output = attn_output.reshape(attn_output.shape[0], attn_output.shape[1], self.num_heads_per_partition * self.head_dim)
attn_output, _ = self.out_proj(attn_output)
return attn_output, None
class CLIPMLP(nn.Module):
def __init__(
self,
config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
self.activation_fn = get_act_fn(config.hidden_act)
self.fc1 = ColumnParallelLinear(config.hidden_size,
config.intermediate_size,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.fc1")
self.fc2 = RowParallelLinear(config.intermediate_size,
config.hidden_size,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.fc2")
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.fc1(hidden_states)
hidden_states = self.activation_fn(hidden_states)
hidden_states, _ = self.fc2(hidden_states)
return hidden_states
class CLIPEncoderLayer(nn.Module):
def __init__(
self,
config: CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__()
self.self_attn = CLIPAttention(
config,
quant_config=quant_config,
prefix=f"{prefix}.self_attn",
)
self.layer_norm1 = nn.LayerNorm(config.hidden_size,
eps=config.layer_norm_eps)
self.mlp = CLIPMLP(config,
quant_config=quant_config,
prefix=f"{prefix}.mlp")
self.layer_norm2 = nn.LayerNorm(config.hidden_size,
eps=config.layer_norm_eps)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
residual = hidden_states
hidden_states = self.layer_norm1(hidden_states)
hidden_states, _ = self.self_attn(hidden_states=hidden_states)
hidden_states = residual + hidden_states
residual = hidden_states
hidden_states = self.layer_norm2(hidden_states)
hidden_states = self.mlp(hidden_states)
hidden_states = residual + hidden_states
return hidden_states
class CLIPEncoder(nn.Module):
"""
Transformer encoder consisting of `config.num_hidden_layers` self
attention layers. Each layer is a [`CLIPEncoderLayer`].
Args:
config: CLIPConfig
"""
def __init__(
self,
config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
if num_hidden_layers_override is None:
num_hidden_layers = config.num_hidden_layers
else:
num_hidden_layers = num_hidden_layers_override
self.layers = nn.ModuleList([
CLIPEncoderLayer(config=config,
quant_config=quant_config,
prefix=f"{prefix}.layers.{layer_idx}")
for layer_idx in range(num_hidden_layers)
])
def forward(
self, inputs_embeds: torch.Tensor, return_all_hidden_states: bool
) -> Union[torch.Tensor, list[torch.Tensor]]:
hidden_states_pool = [inputs_embeds]
hidden_states = inputs_embeds
for encoder_layer in self.layers:
hidden_states = encoder_layer(hidden_states)
if return_all_hidden_states:
hidden_states_pool.append(hidden_states)
# If we have multiple feature sample layers, we return all hidden
# states in order and grab the ones we need by index.
if return_all_hidden_states:
return hidden_states_pool
return [hidden_states]
class CLIPTextTransformer(nn.Module):
def __init__(self,
config: CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None,
*,
num_hidden_layers_override: Optional[int] = None,
prefix: str = ""):
super().__init__()
self.config = config
embed_dim = config.hidden_size
self.embeddings = CLIPTextEmbeddings(config)
self.encoder = CLIPEncoder(config,
quant_config=quant_config,
num_hidden_layers_override=num_hidden_layers_override,
prefix=prefix)
self.final_layer_norm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
# For `pooled_output` computation
self.eos_token_id = config.eos_token_id
# For attention mask, it differs between `flash_attention_2` and other attention implementations
self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2"
def forward(
self,
input_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, BaseModelOutputWithPooling]:
r"""
Returns:
"""
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
if input_ids is None:
raise ValueError("You have to specify input_ids")
input_shape = input_ids.size()
input_ids = input_ids.view(-1, input_shape[-1])
hidden_states = self.embeddings(input_ids=input_ids, position_ids=position_ids)
# CLIP's text model uses causal mask, prepare it here.
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
# causal_attention_mask = _create_4d_causal_attention_mask(
# input_shape, hidden_states.dtype, device=hidden_states.device
# )
# # expand attention_mask
# if attention_mask is not None and not self._use_flash_attention_2:
# raise NotImplementedError("attention_mask is not supported for CLIPTextTransformer")
# # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
# attention_mask = _prepare_4d_attention_mask(attention_mask, hidden_states.dtype)
encoder_outputs = self.encoder(
inputs_embeds=hidden_states,
# attention_mask=attention_mask,
# causal_attention_mask=causal_attention_mask,
# output_attentions=output_attentions,
return_all_hidden_states=output_hidden_states,
# return_dict=return_dict,
)
last_hidden_state = encoder_outputs[-1]
last_hidden_state = self.final_layer_norm(last_hidden_state)
if self.eos_token_id == 2:
# The `eos_token_id` was incorrect before PR #24773: Let's keep what have been done here.
# A CLIP model with such `eos_token_id` in the config can't work correctly with extra new tokens added
# ------------------------------------------------------------
# text_embeds.shape = [batch_size, sequence_length, transformer.width]
# take features from the eot embedding (eot_token is the highest number in each sequence)
# casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14
pooled_output = last_hidden_state[
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
input_ids.to(dtype=torch.int, device=last_hidden_state.device).argmax(dim=-1),
]
else:
# The config gets updated `eos_token_id` from PR #24773 (so the use of exta new tokens is possible)
pooled_output = last_hidden_state[
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
# We need to get the first position of `eos_token_id` value (`pad_token_ids` might equal to `eos_token_id`)
# Note: we assume each sequence (along batch dim.) contains an `eos_token_id` (e.g. prepared by the tokenizer)
(input_ids.to(dtype=torch.int, device=last_hidden_state.device) == self.eos_token_id)
.int()
.argmax(dim=-1),
]
if not return_dict:
return (last_hidden_state, pooled_output) + encoder_outputs[1:]
# return last_hidden_state
return BaseModelOutputWithPooling(
last_hidden_state=last_hidden_state,
pooler_output=pooled_output,
hidden_states=encoder_outputs,
# attentions=encoder_outputs.attentions,
)
class CLIPTextModel(nn.Module):
def __init__(
self,
config: CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
self.text_model = CLIPTextTransformer(
config=config,
quant_config=quant_config,
prefix=prefix)
def forward(
self,
input_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, BaseModelOutputWithPooling]:
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
return self.text_model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=None,
)
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
# Define mapping for stacked parameters
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
]
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
for name, loaded_weight in weights:
# Handle q_proj, k_proj, v_proj -> qkv_proj mapping
for param_name, weight_name, shard_id in stacked_params_mapping:
if weight_name in name:
# Replace the weight name with the parameter name
model_param_name = name.replace(weight_name, param_name)
if model_param_name in params_dict:
param = params_dict[model_param_name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
loaded_params.add(model_param_name)
break
else:
# Use default weight loader for all other parameters
if name in params_dict:
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
class CLIPVisionTransformer(nn.Module):
def __init__(
self,
config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None,
*,
num_hidden_layers_override: Optional[int] = None,
require_post_norm: Optional[bool] = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
embed_dim = config.hidden_size
self.embeddings = CLIPVisionEmbeddings(config)
# NOTE: This typo of "layrnorm" is not fixed on purpose to match
# the original transformers code and name of the model weights.
self.pre_layrnorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
self.encoder = CLIPEncoder(
config=config,
quant_config=quant_config,
num_hidden_layers_override=num_hidden_layers_override,
prefix=f"{prefix}.encoder",
)
num_hidden_layers = config.num_hidden_layers
if len(self.encoder.layers) > config.num_hidden_layers:
raise ValueError(
f"The original encoder only has {num_hidden_layers} "
f"layers, but you requested {len(self.encoder.layers)} layers."
)
# If possible, skip post_layernorm to conserve memory
if require_post_norm is None:
require_post_norm = len(self.encoder.layers) == num_hidden_layers
if require_post_norm:
self.post_layernorm = nn.LayerNorm(embed_dim,
eps=config.layer_norm_eps)
else:
self.post_layernorm = None
def forward(
self,
pixel_values: torch.Tensor,
feature_sample_layers: Optional[list[int]] = None,
) -> torch.Tensor:
hidden_states = self.embeddings(pixel_values)
hidden_states = self.pre_layrnorm(hidden_states)
return_all_hidden_states = feature_sample_layers is not None
# Produces either the last layer output or all of the hidden states,
# depending on if we have feature_sample_layers or not
encoder_outputs = self.encoder(
inputs_embeds=hidden_states,
return_all_hidden_states=return_all_hidden_states)
# Handle post-norm (if applicable) and stacks feature layers if needed
encoder_outputs = resolve_visual_encoder_outputs(
encoder_outputs, feature_sample_layers, self.post_layernorm,
self.config.num_hidden_layers)
return encoder_outputs
class CLIPVisionModel(nn.Module, SupportsQuant):
config_class = CLIPVisionConfig
main_input_name = "pixel_values"
packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]}
def __init__(
self,
config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None,
*,
num_hidden_layers_override: Optional[int] = None,
require_post_norm: Optional[bool] = None,
prefix: str = "",
) -> None:
super().__init__()
self.vision_model = CLIPVisionTransformer(
config=config,
quant_config=quant_config,
num_hidden_layers_override=num_hidden_layers_override,
require_post_norm=require_post_norm,
prefix=f"{prefix}.vision_model")
def forward(
self,
pixel_values: torch.Tensor,
feature_sample_layers: Optional[list[int]] = None,
) -> torch.Tensor:
return self.vision_model(pixel_values, feature_sample_layers)
@property
def device(self):
return next(self.parameters()).device
# (TODO) Add prefix argument for filtering out weights to be loaded
# ref: https://github.com/vllm-project/vllm/pull/7186#discussion_r1734163986
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
]
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
layer_count = len(self.vision_model.encoder.layers)
for name, loaded_weight in weights:
# post_layernorm is not needed in CLIPVisionModel
if (name.startswith("vision_model.post_layernorm")
and self.vision_model.post_layernorm is None):
continue
# omit layers when num_hidden_layers_override is set
if name.startswith("vision_model.encoder.layers"):
layer_idx = int(name.split(".")[3])
if layer_idx >= layer_count:
continue
for (param_name, weight_name, shard_id) in stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
break
else:
param = params_dict[name]
weight_loader = getattr(param, "weight_loader",
default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
-424
View File
@@ -1,424 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from
# https://github.com/huggingface/transformers/blob/v4.28.0/src/transformers/models/llama/modeling_llama.py
# Copyright 2023 The vLLM team.
# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.
#
# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
# and OPT implementations in this library. It has been modified from its
# original forms to accommodate minor architectural differences compared
# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
#
# 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.
"""Inference-only LLaMA model compatible with HuggingFace weights."""
from typing import Any, Dict, Iterable, Optional, Set, Tuple, Type, Union
import torch
from torch import nn
from transformers import LlamaConfig
from transformers.modeling_outputs import BaseModelOutputWithPast
from vllm.attention.layer import MultiHeadAttention
from fastvideo.v1.distributed import get_tensor_model_parallel_world_size
from fastvideo.v1.layers.activation import SiluAndMul
from fastvideo.v1.layers.layernorm import RMSNorm
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
QKVParallelLinear,
RowParallelLinear)
# from vllm.model_executor.layers.quantization import QuantizationConfig
from fastvideo.v1.layers.rotary_embedding import get_rope
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.v1.models.loader.weight_utils import (
default_weight_loader, maybe_remap_kv_scale_name)
from .utils import (extract_layer_index)
class QuantizationConfig:
pass
class LlamaMLP(nn.Module):
def __init__(
self,
hidden_size: int,
intermediate_size: int,
hidden_act: str,
quant_config: Optional[QuantizationConfig] = None,
bias: bool = False,
prefix: str = "",
) -> None:
super().__init__()
self.gate_up_proj = MergedColumnParallelLinear(
input_size=hidden_size,
output_sizes=[intermediate_size] * 2,
# output_size=intermediate_size,
bias=bias,
quant_config=quant_config,
prefix=f"{prefix}.gate_up_proj",
)
self.down_proj = RowParallelLinear(
input_size=intermediate_size,
output_size=hidden_size,
bias=bias,
quant_config=quant_config,
prefix=f"{prefix}.down_proj",
)
if hidden_act != "silu":
raise ValueError(f"Unsupported activation: {hidden_act}. "
"Only silu is supported for now.")
self.act_fn = SiluAndMul()
def forward(self, x):
x, _ = self.gate_up_proj(x)
x = self.act_fn(x)
x, _ = self.down_proj(x)
return x
class LlamaAttention(nn.Module):
def __init__(self,
config: LlamaConfig,
hidden_size: int,
num_heads: int,
num_kv_heads: int,
rope_theta: float = 10000,
rope_scaling: Optional[Dict[str, Any]] = None,
max_position_embeddings: int = 8192,
quant_config: Optional[QuantizationConfig] = None,
bias: bool = False,
bias_o_proj: bool = False,
prefix: str = "") -> None:
super().__init__()
layer_idx = extract_layer_index(prefix)
self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size()
self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size
self.total_num_kv_heads = num_kv_heads
if self.total_num_kv_heads >= tp_size:
# Number of KV heads is greater than TP size, so we partition
# the KV heads across multiple tensor parallel GPUs.
assert self.total_num_kv_heads % tp_size == 0
else:
# Number of KV heads is less than TP size, so we replicate
# the KV heads across multiple tensor parallel GPUs.
assert tp_size % self.total_num_kv_heads == 0
self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)
# MistralConfig has an optional head_dim introduced by Mistral-Nemo
self.head_dim = getattr(config, "head_dim",
self.hidden_size // self.total_num_heads)
# Phi models introduced a partial_rotary_factor parameter in the config
partial_rotary_factor = getattr(config, "partial_rotary_factor", 1)
self.rotary_dim = int(partial_rotary_factor * self.head_dim)
self.q_size = self.num_heads * self.head_dim
self.kv_size = self.num_kv_heads * self.head_dim
self.scaling = self.head_dim**-0.5
self.rope_theta = rope_theta
self.max_position_embeddings = max_position_embeddings
self.qkv_proj = QKVParallelLinear(
hidden_size=hidden_size,
head_size=self.head_dim,
total_num_heads=self.total_num_heads,
total_num_kv_heads=self.total_num_kv_heads,
bias=bias,
quant_config=quant_config,
prefix=f"{prefix}.qkv_proj",
)
self.o_proj = RowParallelLinear(
input_size=self.total_num_heads * self.head_dim,
output_size=hidden_size,
bias=bias_o_proj,
quant_config=quant_config,
prefix=f"{prefix}.o_proj",
)
is_neox_style = True
is_gguf = quant_config and quant_config.get_name() == "gguf"
if is_gguf and config.model_type == "llama":
is_neox_style = False
self.rotary_emb = get_rope(
self.head_dim,
rotary_dim=self.rotary_dim,
max_position=max_position_embeddings,
base=rope_theta,
rope_scaling=rope_scaling,
is_neox_style=is_neox_style,
)
self.attn = MultiHeadAttention(self.num_heads,
self.head_dim,
self.scaling,
self.num_kv_heads)
def forward(
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
) -> torch.Tensor:
qkv, _ = self.qkv_proj(hidden_states)
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
q, k = self.rotary_emb(positions, q, k)
# attn_output = self.attn(q, k, v)
# use flash_attn_func
# TODO (Attn abstraction and backend)
from flash_attn import flash_attn_func
# reshape q, k, v to (batch_size, seq_len, num_heads, head_dim)
batch_size = q.shape[0]
seq_len = q.shape[1]
q = q.reshape(batch_size, seq_len, self.num_heads, self.head_dim)
k = k.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim)
v = v.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim)
# import pdb; pdb.set_trace()
attn_output = flash_attn_func(q, k, v, softmax_scale=self.scaling, causal=True)
attn_output = attn_output.reshape(batch_size, seq_len, self.num_heads * self.head_dim)
output, _ = self.o_proj(attn_output)
return output
class LlamaDecoderLayer(nn.Module):
def __init__(
self,
config: LlamaConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__()
self.hidden_size = config.hidden_size
rope_theta = getattr(config, "rope_theta", 10000)
rope_scaling = getattr(config, "rope_scaling", None)
if rope_scaling is not None and getattr(
config, "original_max_position_embeddings", None):
rope_scaling["original_max_position_embeddings"] = (
config.original_max_position_embeddings)
max_position_embeddings = getattr(config, "max_position_embeddings",
8192)
# Support abacusai/Smaug-72B-v0.1 with attention_bias
# Support internlm/internlm-7b with bias
attention_bias = getattr(config, "attention_bias", False) or getattr(
config, "bias", False)
bias_o_proj = attention_bias
# support internlm/internlm3-8b with qkv_bias
if hasattr(config, 'qkv_bias'):
attention_bias = config.qkv_bias
self.self_attn = LlamaAttention(
config=config,
hidden_size=self.hidden_size,
num_heads=config.num_attention_heads,
num_kv_heads=getattr(config, "num_key_value_heads",
config.num_attention_heads),
rope_theta=rope_theta,
rope_scaling=rope_scaling,
max_position_embeddings=max_position_embeddings,
quant_config=quant_config,
bias=attention_bias,
bias_o_proj=bias_o_proj,
prefix=f"{prefix}.self_attn",
)
self.mlp = LlamaMLP(
hidden_size=self.hidden_size,
intermediate_size=config.intermediate_size,
hidden_act=config.hidden_act,
quant_config=quant_config,
bias=getattr(config, "mlp_bias", False),
prefix=f"{prefix}.mlp",
)
self.input_layernorm = RMSNorm(config.hidden_size,
eps=config.rms_norm_eps)
self.post_attention_layernorm = RMSNorm(config.hidden_size,
eps=config.rms_norm_eps)
def forward(
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
residual: Optional[torch.Tensor],
) -> Tuple[torch.Tensor, torch.Tensor]:
# Self Attention
if residual is None:
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
else:
hidden_states, residual = self.input_layernorm(
hidden_states, residual)
hidden_states = self.self_attn(positions=positions,
hidden_states=hidden_states)
# Fully Connected
hidden_states, residual = self.post_attention_layernorm(
hidden_states, residual)
hidden_states = self.mlp(hidden_states)
return hidden_states, residual
class LlamaModel(nn.Module):
def __init__(self,
config: LlamaConfig,
prefix: str = "",
layer_type: Type[LlamaDecoderLayer] = LlamaDecoderLayer):
super().__init__()
quant_config = None
lora_config = None
self.config = config
self.quant_config = quant_config
lora_vocab = (lora_config.lora_extra_vocab_size *
(lora_config.max_loras or 1)) if lora_config else 0
self.vocab_size = config.vocab_size + lora_vocab
self.org_vocab_size = config.vocab_size
self.embed_tokens = VocabParallelEmbedding(
self.vocab_size,
config.hidden_size,
org_num_embeddings=config.vocab_size,
quant_config=quant_config,
)
self.layers = nn.ModuleList([
layer_type(
config=config,
quant_config=quant_config,
prefix=f"{prefix}.layers.{i}"
) for i in range(config.num_hidden_layers)
])
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.embed_tokens(input_ids)
def forward(
self,
input_ids: Optional[torch.Tensor],
positions: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.Tensor] = None,
output_hidden_states: Optional[bool] = None,
) -> torch.Tensor:
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
if inputs_embeds is not None:
hidden_states = inputs_embeds
else:
hidden_states = self.get_input_embeddings(input_ids)
residual = None
if positions is None:
positions = torch.arange(
0, hidden_states.shape[1], device=hidden_states.device
).unsqueeze(0)
all_hidden_states = () if output_hidden_states else None
for layer in self.layers:
if output_hidden_states:
all_hidden_states += (hidden_states,)
hidden_states, residual = layer(positions, hidden_states, residual)
hidden_states, _ = self.norm(hidden_states, residual)
# add hidden states from the last decoder layer
if output_hidden_states:
all_hidden_states += (hidden_states,)
# TODO(will): maybe unify the output format with other models and use
# our own class
output = BaseModelOutputWithPast(
last_hidden_state=hidden_states,
# past_key_values=past_key_values if use_cache else None,
hidden_states=all_hidden_states,
# attentions=all_self_attns,
)
return output
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
(".qkv_proj", ".q_proj", "q"),
(".qkv_proj", ".k_proj", "k"),
(".qkv_proj", ".v_proj", "v"),
(".gate_up_proj", ".gate_proj", 0),
(".gate_up_proj", ".up_proj", 1),
]
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
for name, loaded_weight in weights:
if "rotary_emb.inv_freq" in name:
continue
if ("rotary_emb.cos_cached" in name
or "rotary_emb.sin_cached" in name):
# Models trained using ColossalAI may include these tensors in
# the checkpoint. Skip them.
continue
if (self.quant_config is not None and
(scale_name := self.quant_config.get_cache_scale(name))):
# Loading kv cache quantization scales
param = params_dict[scale_name]
weight_loader = getattr(param, "weight_loader",
default_weight_loader)
loaded_weight = (loaded_weight if loaded_weight.dim() == 0 else
loaded_weight[0])
weight_loader(param, loaded_weight)
loaded_params.add(scale_name)
continue
if "scale" in name:
# Remapping the name of FP8 kv-scale.
name = maybe_remap_kv_scale_name(name, params_dict)
if name is None:
continue
for param_name, weight_name, shard_id in stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
# Skip loading extra bias for GPTQ models.
if name.endswith(".bias") and name not in params_dict:
continue
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
break
else:
# Skip loading extra bias for GPTQ models.
if name.endswith(".bias") and name not in params_dict:
continue
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader",
default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
-787
View File
@@ -1,787 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Derived from T5 implementation posted on HuggingFace; license below:
#
# coding=utf-8
# Copyright 2018 Mesh TensorFlow authors, T5 Authors and HuggingFace Inc. team.
#
# 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.
"""PyTorch T5 model."""
import math
import re
from typing import Iterable, List, Optional, Set, Tuple
import torch
from torch import nn
from transformers import T5Config
# TODO best way to handle xformers imports?
from xformers.ops.fmha.attn_bias import LowerTriangularMaskWithTensorBias
# TODO func should be in backend interface
from vllm.attention.backends.xformers import (XFormersMetadata, _get_attn_bias,
_set_attn_bias)
from vllm.attention.layer import Attention, AttentionMetadata, AttentionType
from vllm.config import CacheConfig, VllmConfig
from vllm.distributed import get_tensor_model_parallel_world_size
from vllm.model_executor.layers.activation import get_act_fn
from vllm.model_executor.layers.linear import (ColumnParallelLinear,
QKVParallelLinear,
RowParallelLinear)
from vllm.model_executor.layers.logits_processor import LogitsProcessor
from vllm.model_executor.layers.quantization.base_config import (
QuantizationConfig)
from vllm.model_executor.layers.sampler import SamplerOutput, get_sampler
from vllm.model_executor.layers.vocab_parallel_embedding import (
ParallelLMHead, VocabParallelEmbedding)
from vllm.model_executor.model_loader.weight_utils import default_weight_loader
from vllm.model_executor.sampling_metadata import SamplingMetadata
from vllm.sequence import IntermediateTensors
from .utils import maybe_prefix
class T5LayerNorm(nn.Module):
def __init__(self, hidden_size, eps=1e-6):
"""
Construct a layernorm module in the T5 style.
No bias and no subtraction of mean.
"""
super().__init__()
self.weight = nn.Parameter(torch.ones(hidden_size))
self.variance_epsilon = eps
def forward(self, hidden_states) -> torch.Tensor:
# T5 uses a layer_norm which only scales and doesn't shift, which is
# also known as Root Mean Square Layer Normalization
# https://arxiv.org/abs/1910.07467 thus variance is calculated w/o mean
# and there is no bias. Additionally we want to make sure that the
# accumulation for half-precision inputs is done in fp32.
# TODO (rmns norm ops)
variance = hidden_states.to(torch.float32).pow(2).mean(-1,
keepdim=True)
hidden_states = hidden_states * torch.rsqrt(variance +
self.variance_epsilon)
# convert into half-precision if necessary
if self.weight.dtype in [torch.float16, torch.bfloat16]:
hidden_states = hidden_states.to(self.weight.dtype)
return self.weight * hidden_states
class T5DenseActDense(nn.Module):
def __init__(self,
config: T5Config,
quant_config: Optional[QuantizationConfig] = None):
super().__init__()
self.wi = ColumnParallelLinear(config.d_model, config.d_ff, bias=False)
self.wo = RowParallelLinear(config.d_ff,
config.d_model,
bias=False,
quant_config=quant_config)
self.act = get_act_fn(config.dense_act_fn)
def forward(self, hidden_states) -> torch.Tensor:
hidden_states, _ = self.wi(hidden_states)
hidden_states = self.act(hidden_states)
# if (
# isinstance(self.wo.weight, torch.Tensor)
# and hidden_states.dtype != self.wo.weight.dtype
# and self.wo.weight.dtype != torch.int8
# ):
# hidden_states = hidden_states.to(self.wo.weight.dtype)
hidden_states, _ = self.wo(hidden_states)
return hidden_states
class T5DenseGatedActDense(nn.Module):
def __init__(self,
config: T5Config,
quant_config: Optional[QuantizationConfig] = None):
super().__init__()
self.wi_0 = ColumnParallelLinear(config.d_model,
config.d_ff,
bias=False,
quant_config=quant_config)
self.wi_1 = ColumnParallelLinear(config.d_model,
config.d_ff,
bias=False,
quant_config=quant_config)
# Should not run in fp16 unless mixed-precision is used,
# see https://github.com/huggingface/transformers/issues/20287.
self.wo = RowParallelLinear(config.d_ff,
config.d_model,
bias=False,
quant_config=quant_config)
self.act = get_act_fn(config.dense_act_fn)
def forward(self, hidden_states) -> torch.Tensor:
hidden_gelu = self.act(self.wi_0(hidden_states)[0])
hidden_linear, _ = self.wi_1(hidden_states)
hidden_states = hidden_gelu * hidden_linear
hidden_states, _ = self.wo(hidden_states)
return hidden_states
class T5LayerFF(nn.Module):
def __init__(self,
config: T5Config,
quant_config: Optional[QuantizationConfig] = None):
super().__init__()
if config.is_gated_act:
self.DenseReluDense = T5DenseGatedActDense(
config, quant_config=quant_config)
else:
self.DenseReluDense = T5DenseActDense(config,
quant_config=quant_config)
self.layer_norm = T5LayerNorm(config.d_model,
eps=config.layer_norm_epsilon)
def forward(self, hidden_states) -> torch.Tensor:
forwarded_states = self.layer_norm(hidden_states)
forwarded_states = self.DenseReluDense(forwarded_states)
hidden_states = hidden_states + forwarded_states
return hidden_states
class T5Attention(nn.Module):
def __init__(self,
config: T5Config,
attn_type: AttentionType,
has_relative_attention_bias=False,
cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
super().__init__()
self.attn_type = attn_type
# Cross-attention has no relative pos encoding anyway
self.is_decoder = attn_type == AttentionType.DECODER
self.has_relative_attention_bias = has_relative_attention_bias
self.relative_attention_num_buckets = \
config.relative_attention_num_buckets
self.relative_attention_max_distance = \
config.relative_attention_max_distance
self.d_model = config.d_model
self.key_value_proj_dim = config.d_kv
assert cache_config
# Alternatively we can get it from kv_cache size in fwd.
self.block_size = cache_config.block_size
# Partition heads across multiple tensor parallel GPUs.
tp_world_size = get_tensor_model_parallel_world_size()
assert config.num_heads % tp_world_size == 0
self.n_heads = config.num_heads // tp_world_size
self.inner_dim = self.n_heads * self.key_value_proj_dim
# No GQA in t5.
self.n_kv_heads = self.n_heads
self.qkv_proj = QKVParallelLinear(self.d_model,
self.d_model // self.n_heads,
self.n_heads,
self.n_kv_heads,
bias=False,
quant_config=quant_config)
# NOTE (NickLucche) T5 employs a scaled weight initialization scheme
# instead of scaling attention scores directly.
self.attn = Attention(self.n_heads,
config.d_kv,
1.0,
cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.attn",
attn_type=self.attn_type)
# Only the first SelfAttention block in encoder decoder has this
# embedding layer, the others reuse its output.
if self.has_relative_attention_bias:
self.relative_attention_bias = \
VocabParallelEmbedding(self.relative_attention_num_buckets,
self.n_heads,
org_num_embeddings=\
self.relative_attention_num_buckets,
quant_config=quant_config)
self.out_proj = RowParallelLinear(
self.inner_dim,
self.d_model,
bias=False,
quant_config=quant_config,
)
@staticmethod
def _relative_position_bucket(relative_position,
bidirectional=True,
num_buckets=32,
max_distance=128):
"""
Adapted from Mesh Tensorflow:
https://github.com/tensorflow/mesh/blob/0cb87fe07da627bf0b7e60475d59f95ed6b5be3d/mesh_tensorflow/transformer/transformer_layers.py#L593
Translate relative position to a bucket number for relative attention.
The relative position is defined as memory_position - query_position,
i.e. the distance in tokens from the attending position to the
attended-to position. If bidirectional=False, then positive relative
positions are invalid. We use smaller buckets for small absolute
relative_position and larger buckets for larger absolute
relative_positions. All relative positions >=max_distance map to the
same bucket. All relative positions <=-max_distance map to the same
bucket. This should allow for more graceful generalization to longer
sequences than the model has been trained on
Args:
relative_position: an int32 Tensor
bidirectional: a boolean - whether the attention is bidirectional
num_buckets: an integer
max_distance: an integer
Returns:
a Tensor with the same shape as relative_position, containing int32
values in the range [0, num_buckets)
"""# noqa: E501
relative_buckets = 0
if bidirectional:
num_buckets //= 2
relative_buckets += (relative_position > 0).to(
torch.long) * num_buckets
relative_position = torch.abs(relative_position)
else:
relative_position = -torch.min(relative_position,
torch.zeros_like(relative_position))
# now relative_position is in the range [0, inf)
# half of the buckets are for exact increments in positions
max_exact = num_buckets // 2
is_small = relative_position < max_exact
# The other half of the buckets are for logarithmically bigger bins
# in positions up to max_distance
relative_position_if_large = max_exact + (
torch.log(relative_position.float() / max_exact) /
math.log(max_distance / max_exact) *
(num_buckets - max_exact)).to(torch.long)
relative_position_if_large = torch.min(
relative_position_if_large,
torch.full_like(relative_position_if_large, num_buckets - 1))
relative_buckets += torch.where(is_small, relative_position,
relative_position_if_large)
return relative_buckets
def compute_bias(self,
query_length,
key_length,
device=None) -> torch.Tensor:
"""Compute binned relative position bias"""
# TODO possible tp issue?
if device is None:
device = self.relative_attention_bias.weight.device
context_position = torch.arange(query_length,
dtype=torch.long,
device=device)[:, None]
memory_position = torch.arange(key_length,
dtype=torch.long,
device=device)[None, :]
# max_seq_len, nh
relative_position = memory_position - context_position
relative_position_bucket = self._relative_position_bucket(
relative_position, # shape (query_length, key_length)
bidirectional=(not self.is_decoder),
num_buckets=self.relative_attention_num_buckets,
max_distance=self.relative_attention_max_distance,
)
values = self.relative_attention_bias(
relative_position_bucket
) # shape (query_length, key_length, num_heads)
x = values.permute([2, 0, 1]).unsqueeze(
0) # shape (1, num_heads, query_length, key_length)
return x
def forward(
self,
hidden_states: torch.Tensor, # (num_tokens, d_model)
kv_cache: torch.Tensor,
attn_metadata: AttentionMetadata,
encoder_hidden_states: Optional[torch.Tensor] = None,
) -> torch.Tensor:
# TODO auto-selection of xformers backend when t5 is detected
assert isinstance(attn_metadata, XFormersMetadata)
num_seqs = len(
attn_metadata.seq_lens) if attn_metadata.seq_lens else len(
attn_metadata.encoder_seq_lens)
qkv, _ = self.qkv_proj(hidden_states)
# Projection of 'own' hidden state (self-attention). No GQA here.
q, k, v = qkv.split(self.inner_dim, dim=-1)
# NOTE (NickLucche) Attn bias is computed once per encoder or decoder
# forward, on the first call to T5Attention.forward. Subsequent
# *self-attention* layers will reuse it.
attn_bias = _get_attn_bias(attn_metadata, self.attn_type)
if self.attn_type == AttentionType.ENCODER_DECODER:
# Projection of encoder's hidden states, cross-attention.
if encoder_hidden_states is None:
# Decode phase, kv already cached
assert attn_metadata.num_prefills == 0
k = None
v = None
else:
assert attn_metadata.num_prefills > 0
# Prefill phase (first decoder forward), caching kv
qkv_enc, _ = self.qkv_proj(encoder_hidden_states)
_, k, v = qkv_enc.split(self.inner_dim, dim=-1)
# No custom attention bias must be set when running cross attn.
assert attn_bias is None
# Not compatible with CP here (as all encoder-decoder models),
# as it assumes homogeneous batch (prefills or decodes).
elif self.has_relative_attention_bias:
assert attn_bias is None # to be recomputed
# Self-attention. Compute T5 relative positional encoding.
# The bias term is computed on longest sequence in batch. Biases
# for shorter sequences are slices of the longest.
# TODO xformers-specific code.
align_to = 8
# bias expected shape: (num_seqs, NH, L, L_pad) for prefill,
# (num_seqs, NH, 1, L_pad) for decodes.
if self.attn_type == AttentionType.ENCODER:
# Encoder prefill stage, uses xFormers, hence sequence
# padding/alignment to 8 is required.
seq_len = attn_metadata.max_encoder_seq_len
padded_seq_len = (seq_len + align_to -
1) // align_to * align_to
# TODO (NickLucche) avoid extra copy on repeat,
# provide multiple slices of same memory
position_bias = self.compute_bias(seq_len,
padded_seq_len).repeat(
num_seqs, 1, 1, 1)
# xFormers expects a list of biases, one matrix per sequence.
# As each sequence gets its own bias, no masking is required.
attn_bias = [
p[None, :, :sq, :sq] for p, sq in zip(
position_bias, attn_metadata.encoder_seq_lens)
]
elif attn_metadata.prefill_metadata:
# Decoder prefill stage, uses xFormers, hence sequence
# padding/alignment to 8 is required. First decoder step,
# seq_len is usually 1, but one can prepend different start
# tokens prior to generation.
seq_len = attn_metadata.max_prefill_seq_len
# ->align
padded_seq_len = (seq_len + align_to -
1) // align_to * align_to
position_bias = self.compute_bias(seq_len,
padded_seq_len).repeat(
num_seqs, 1, 1, 1)
# Causal mask for prefill.
attn_bias = [
LowerTriangularMaskWithTensorBias(pb[None, :, :sq, :sq])
for pb, sq in zip(position_bias, attn_metadata.seq_lens)
]
else:
# Decoder decoding stage, uses PagedAttention, hence sequence
# padding/alignment to `block_size` is required. Expected
# number of queries is always 1 (MQA not supported).
seq_len = attn_metadata.max_decode_seq_len
block_aligned_seq_len = (seq_len + self.block_size - 1
) // self.block_size * self.block_size
# TODO bf16 bias support in PagedAttention.
position_bias = self.compute_bias(
seq_len, block_aligned_seq_len).float()
# Bias for the last query, the one at current decoding step.
position_bias = position_bias[:, :, -1:, :].repeat(
num_seqs, 1, 1, 1)
# No explicit masking required, this is done inside the
# paged attention kernel based on the sequence length.
attn_bias = [position_bias]
# NOTE Assign bias term on metadata based on attn type:
# ENCODER->`encoder_attn_bias`, DECODER->`attn_bias`.
_set_attn_bias(attn_metadata, attn_bias, self.attn_type)
elif not self.has_relative_attention_bias:
# Encoder/Decoder Self-Attention Layer, attn bias already cached.
assert attn_bias is not None
attn_output = self.attn(q, k, v, kv_cache, attn_metadata)
output, _ = self.out_proj(attn_output)
return output
class T5LayerSelfAttention(nn.Module):
def __init__(
self,
config,
has_relative_attention_bias=False,
cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
self.SelfAttention = T5Attention(
config,
AttentionType.DECODER
if "decoder" in prefix else AttentionType.ENCODER,
has_relative_attention_bias=has_relative_attention_bias,
cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.SelfAttention")
self.layer_norm = T5LayerNorm(config.d_model,
eps=config.layer_norm_epsilon)
def forward(
self,
hidden_states: torch.Tensor,
kv_cache: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
normed_hidden_states = self.layer_norm(hidden_states)
attention_output = self.SelfAttention(
hidden_states=normed_hidden_states,
kv_cache=kv_cache,
attn_metadata=attn_metadata,
encoder_hidden_states=None,
)
hidden_states = hidden_states + attention_output
return hidden_states
class T5LayerCrossAttention(nn.Module):
def __init__(self,
config,
cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
super().__init__()
self.EncDecAttention = T5Attention(config,
AttentionType.ENCODER_DECODER,
has_relative_attention_bias=False,
cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.EncDecAttention")
self.layer_norm = T5LayerNorm(config.d_model,
eps=config.layer_norm_epsilon)
def forward(
self,
hidden_states: torch.Tensor,
kv_cache: torch.Tensor,
attn_metadata: AttentionMetadata,
encoder_hidden_states: Optional[torch.Tensor] = None,
) -> torch.Tensor:
normed_hidden_states = self.layer_norm(hidden_states)
attention_output = self.EncDecAttention(
hidden_states=normed_hidden_states,
kv_cache=kv_cache,
attn_metadata=attn_metadata,
encoder_hidden_states=encoder_hidden_states,
)
hidden_states = hidden_states + attention_output
return hidden_states
class T5Block(nn.Module):
def __init__(self,
config: T5Config,
is_decoder: bool,
has_relative_attention_bias=False,
cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
super().__init__()
self.is_decoder = is_decoder
self.self_attn = T5LayerSelfAttention(
config,
has_relative_attention_bias=has_relative_attention_bias,
cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.self_attn")
if self.is_decoder:
self.cross_attn = T5LayerCrossAttention(
config,
cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.cross_attn")
self.ffn = T5LayerFF(config, quant_config=quant_config)
def forward(
self,
hidden_states: torch.Tensor,
kv_cache: torch.Tensor,
attn_metadata: AttentionMetadata,
encoder_hidden_states: Optional[torch.Tensor] = None,
) -> torch.Tensor:
hidden_states = self.self_attn(
hidden_states=hidden_states,
kv_cache=kv_cache,
attn_metadata=attn_metadata,
)
if self.is_decoder:
hidden_states = self.cross_attn(
hidden_states=hidden_states,
kv_cache=kv_cache,
attn_metadata=attn_metadata,
encoder_hidden_states=encoder_hidden_states,
)
# Apply Feed Forward layer
hidden_states = self.ffn(hidden_states)
return hidden_states
class T5Stack(nn.Module):
def __init__(self,
config: T5Config,
is_decoder: bool,
n_layers: int,
embed_tokens=None,
cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
super().__init__()
self.embed_tokens = embed_tokens
# Only the first block has relative positional encoding.
self.blocks = nn.ModuleList([
T5Block(config,
is_decoder=is_decoder,
has_relative_attention_bias=i == 0,
cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.blocks.{i}") for i in range(n_layers)
])
self.final_layer_norm = T5LayerNorm(config.d_model,
eps=config.layer_norm_epsilon)
def forward(
self,
input_ids: torch.Tensor,
kv_caches: List[torch.Tensor],
attn_metadata: AttentionMetadata,
encoder_hidden_states: Optional[torch.Tensor] = None
) -> torch.Tensor:
hidden_states = self.embed_tokens(input_ids)
for idx, block in enumerate(self.blocks):
hidden_states = block(
hidden_states=hidden_states,
kv_cache=kv_caches[idx],
attn_metadata=attn_metadata,
encoder_hidden_states=encoder_hidden_states,
)
hidden_states = self.final_layer_norm(hidden_states)
return hidden_states
class T5Model(nn.Module):
_tied_weights_keys = [
"encoder.embed_tokens.weight", "decoder.embed_tokens.weight"
]
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
super().__init__()
config: T5Config = vllm_config.model_config.hf_config
cache_config = vllm_config.cache_config
quant_config = vllm_config.quant_config
lora_config = vllm_config.lora_config
lora_vocab = (lora_config.lora_extra_vocab_size *
(lora_config.max_loras or 1)) if lora_config else 0
self.vocab_size = config.vocab_size + lora_vocab
self.padding_idx = config.pad_token_id
self.shared = VocabParallelEmbedding(
config.vocab_size,
config.d_model,
org_num_embeddings=config.vocab_size)
self.encoder = T5Stack(config,
False,
config.num_layers,
self.shared,
cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.encoder")
self.decoder = T5Stack(config,
True,
config.num_decoder_layers,
self.shared,
cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.decoder")
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.shared(input_ids)
def forward(self, input_ids: torch.Tensor, encoder_input_ids: torch.Tensor,
kv_caches: List[torch.Tensor],
attn_metadata: AttentionMetadata) -> torch.Tensor:
encoder_hidden_states = None
if encoder_input_ids.numel() > 0:
# Run encoder attention if a non-zero number of encoder tokens
# are provided as input: on a regular generate call, the encoder
# runs once, on the prompt. Subsequent decoder calls reuse output
# `encoder_hidden_states`.
encoder_hidden_states = self.encoder(input_ids=encoder_input_ids,
kv_caches=kv_caches,
attn_metadata=attn_metadata)
# Clear attention bias state.
attn_metadata.attn_bias = None
attn_metadata.encoder_attn_bias = None
attn_metadata.cross_attn_bias = None
decoder_outputs = self.decoder(
input_ids=input_ids,
encoder_hidden_states=encoder_hidden_states,
kv_caches=kv_caches,
attn_metadata=attn_metadata)
# When capturing CUDA Graph
attn_metadata.attn_bias = None
attn_metadata.encoder_attn_bias = None
attn_metadata.cross_attn_bias = None
return decoder_outputs
class T5ForConditionalGeneration(nn.Module):
_keys_to_ignore_on_load_unexpected = [
"decoder.block.0.layer.1.EncDecAttention.relative_attention_bias.weight",
]
_tied_weights_keys = [
"encoder.embed_tokens.weight", "decoder.embed_tokens.weight",
"lm_head.weight"
]
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
super().__init__()
config: T5Config = vllm_config.model_config.hf_config
self.model_dim = config.d_model
self.config = config
self.unpadded_vocab_size = config.vocab_size
if lora_config := vllm_config.lora_config:
self.unpadded_vocab_size += lora_config.lora_extra_vocab_size
self.model = T5Model(vllm_config=vllm_config,
prefix=maybe_prefix(prefix, "model"))
# Although not in config, this is the default for hf models.
if self.config.tie_word_embeddings:
self.lm_head = self.model.shared
# in transformers this is smt more explicit, as in (after load)
# self.lm_head.weight = self.model.shared.weight
else:
self.lm_head = ParallelLMHead(self.unpadded_vocab_size,
config.d_model,
org_num_embeddings=config.vocab_size)
self.logits_processor = LogitsProcessor(self.unpadded_vocab_size,
config.vocab_size)
self.sampler = get_sampler()
def compute_logits(
self,
hidden_states: torch.Tensor,
sampling_metadata: SamplingMetadata,
) -> Optional[torch.Tensor]:
if self.config.tie_word_embeddings:
# Rescale output before projecting on vocab
# See https://github.com/tensorflow/mesh/blob/fa19d69eafc9a482aff0b59ddd96b025c0cb207d/mesh_tensorflow/transformer/transformer.py#L586 # noqa: E501
hidden_states = hidden_states * (self.model_dim**-0.5)
logits = self.logits_processor(self.lm_head, hidden_states,
sampling_metadata)
return logits
def sample(
self,
logits: Optional[torch.Tensor],
sampling_metadata: SamplingMetadata,
) -> Optional[SamplerOutput]:
next_tokens = self.sampler(logits, sampling_metadata)
return next_tokens
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.model.shared(input_ids)
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
kv_caches: List[torch.Tensor],
attn_metadata: AttentionMetadata,
intermediate_tensors: Optional[IntermediateTensors] = None,
*,
encoder_input_ids: torch.Tensor,
encoder_positions: torch.Tensor,
**kwargs,
) -> torch.Tensor:
return self.model(input_ids, encoder_input_ids, kv_caches,
attn_metadata)
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
model_params_dict = dict(self.named_parameters(remove_duplicate=False))
loaded_params: Set[str] = set()
renamed_reg = [
(re.compile(r'block\.(\d+)\.layer\.0'), r'blocks.\1.self_attn'),
(re.compile(r'decoder.block\.(\d+)\.layer\.1'),
r'decoder.blocks.\1.cross_attn'),
(re.compile(r'decoder.block\.(\d+)\.layer\.2'),
r'decoder.blocks.\1.ffn'),
# encoder has no cross-attn, but rather self-attention+ffn.
(re.compile(r'encoder.block\.(\d+)\.layer\.1'),
r'encoder.blocks.\1.ffn'),
(re.compile(r'\.o\.'), r'.out_proj.'),
]
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
(".qkv_proj.", ".q.", "q"),
(".qkv_proj.", ".k.", "k"),
(".qkv_proj.", ".v.", "v")
]
for name, loaded_weight in weights:
# No relative position attn bias on cross attention.
if name in self._keys_to_ignore_on_load_unexpected:
continue
# Handle some renaming
for reg in renamed_reg:
name = re.sub(*reg, name)
top_module, _ = name.split('.', 1)
if top_module != 'lm_head':
name = f"model.{name}"
# Split q/k/v layers to unified QKVParallelLinear
for (param_name, weight_name, shard_id) in stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
param = model_params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
break
else:
# Not a q/k/v layer.
param = model_params_dict[name]
weight_loader = getattr(param, "weight_loader",
default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
-22
View File
@@ -1,22 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from typing import List
def extract_layer_index(layer_name: str) -> int:
"""
Extract the layer index from the module name.
Examples:
- "encoder.layers.0" -> 0
- "encoder.layers.1.self_attn" -> 1
- "2.self_attn" -> 2
- "model.encoder.layers.0.sub.1" -> ValueError
"""
subnames = layer_name.split(".")
int_vals: List[int] = []
for subname in subnames:
try:
int_vals.append(int(subname))
except ValueError:
continue
assert len(int_vals) == 1, (f"layer name {layer_name} should"
" only contain one integer")
return int_vals[0]
-150
View File
@@ -1,150 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import Final, Generic, Optional, Protocol, TypeVar, Union
import torch
from transformers import PretrainedConfig
import fastvideo.v1.envs as envs
from vllm.attention.selector import (backend_name_to_enum,
get_global_forced_attn_backend)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.platforms import _Backend, current_platform
logger = init_logger(__name__)
_C = TypeVar("_C", bound=PretrainedConfig)
class VisionEncoderInfo(ABC, Generic[_C]):
def __init__(self, vision_config: _C) -> None:
super().__init__()
self.vision_config = vision_config
@abstractmethod
def get_num_image_tokens(
self,
*,
image_width: int,
image_height: int,
) -> int:
raise NotImplementedError
@abstractmethod
def get_max_image_tokens(self) -> int:
raise NotImplementedError
@abstractmethod
def get_image_size(self) -> int:
raise NotImplementedError
@abstractmethod
def get_patch_size(self) -> int:
raise NotImplementedError
@abstractmethod
def get_patch_grid_length(self) -> int:
raise NotImplementedError
class VisionLanguageConfig(Protocol):
vision_config: Final[PretrainedConfig]
def get_vision_encoder_info(
hf_config: VisionLanguageConfig) -> VisionEncoderInfo:
# Avoid circular imports
from .clip import CLIPEncoderInfo, CLIPVisionConfig
from .pixtral import PixtralHFEncoderInfo, PixtralVisionConfig
from .siglip import SiglipEncoderInfo, SiglipVisionConfig
vision_config = hf_config.vision_config
if isinstance(vision_config, CLIPVisionConfig):
return CLIPEncoderInfo(vision_config)
if isinstance(vision_config, PixtralVisionConfig):
return PixtralHFEncoderInfo(vision_config)
if isinstance(vision_config, SiglipVisionConfig):
return SiglipEncoderInfo(vision_config)
msg = f"Unsupported vision config: {type(vision_config)}"
raise NotImplementedError(msg)
def get_vit_attn_backend(support_fa: bool = False) -> _Backend:
"""
Get the available attention backend for Vision Transformer.
"""
# TODO(Isotr0py): Remove `support_fa` after support FA for all ViTs attn.
selected_backend: Optional[_Backend] = get_global_forced_attn_backend()
if selected_backend is None:
backend_by_env_var: Optional[str] = envs.VLLM_ATTENTION_BACKEND
if backend_by_env_var is not None:
selected_backend = backend_name_to_enum(backend_by_env_var)
if selected_backend is None:
if current_platform.is_cuda():
device_available = current_platform.has_device_capability(80)
if device_available and support_fa:
from transformers.utils import is_flash_attn_2_available
if is_flash_attn_2_available():
selected_backend = _Backend.FLASH_ATTN
else:
logger.warning_once(
"Current `vllm-flash-attn` has a bug inside vision "
"module, so we use xformers backend instead. You can "
"run `pip install flash-attn` to use flash-attention "
"backend.")
selected_backend = _Backend.XFORMERS
else:
# For Volta and Turing GPUs, use xformers instead.
selected_backend = _Backend.XFORMERS
else:
# Default to torch SDPA for other non-GPU platforms.
selected_backend = _Backend.TORCH_SDPA
return selected_backend
def resolve_visual_encoder_outputs(
encoder_outputs: Union[torch.Tensor, list[torch.Tensor]],
feature_sample_layers: Optional[list[int]],
post_layer_norm: Optional[torch.nn.LayerNorm],
max_possible_layers: int,
) -> torch.Tensor:
"""Given the outputs a visual encoder module that may correspond to the
output of the last layer, or a list of hidden states to be stacked,
handle post normalization and resolve it into a single output tensor.
Args:
encoder_outputs: Output of encoder's last layer or all hidden states.
feature_sample_layers: Optional layer indices to grab from the encoder
outputs; if provided, encoder outputs must be a list.
post_layer_norm: Post norm to apply to the output of the encoder.
max_possible_layers: Total layers in the fully loaded visual encoder.
"""
if feature_sample_layers is None:
if post_layer_norm is not None:
return post_layer_norm(encoder_outputs)
return encoder_outputs
# Get the hidden states corresponding to the layer indices.
# Negative values are relative to the full visual encoder,
# so offset them depending on how many layers were loaded.
# NOTE: this assumes that encoder_outputs is a list containing
# the inputs to the visual encoder, followed by the hidden states
# of each layer.
num_loaded_layers = len(encoder_outputs) - 1
offset = max_possible_layers - num_loaded_layers
hs_pool = [
encoder_outputs[layer_idx]
if layer_idx >= 0 else encoder_outputs[layer_idx + offset]
for layer_idx in feature_sample_layers
]
# Apply post-norm on the final hidden state if we are using it
uses_last_layer = feature_sample_layers[-1] in (len(hs_pool) - 1, -1)
if post_layer_norm is not None and uses_last_layer:
hs_pool[-1] = post_layer_norm(encoder_outputs)
return torch.cat(hs_pool, dim=-1)
-259
View File
@@ -1,259 +0,0 @@
# Copyright 2023-2024 SGLang Team
# 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.
# ==============================================================================
"""Utilities for Huggingface Transformers."""
import contextlib
import os
import warnings
from pathlib import Path
from typing import Dict, Optional, Type, Union, Any
import json
from huggingface_hub import snapshot_download
from transformers import (
AutoConfig,
AutoProcessor,
AutoTokenizer,
PretrainedConfig,
PreTrainedTokenizer,
PreTrainedTokenizerFast,
)
from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_NAMES
# from fastvideo.v1.models.configs import ChatGLMConfig, DbrxConfig, ExaoneConfig, Qwen2_5_VLConfig
_CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
# ChatGLMConfig.model_type: ChatGLMConfig,
# DbrxConfig.model_type: DbrxConfig,
# ExaoneConfig.model_type: ExaoneConfig,
# Qwen2_5_VLConfig.model_type: Qwen2_5_VLConfig,
}
for name, cls in _CONFIG_REGISTRY.items():
with contextlib.suppress(ValueError):
AutoConfig.register(name, cls)
def download_from_hf(model_path: str):
if os.path.exists(model_path):
return model_path
return snapshot_download(model_path, allow_patterns=["*.json", "*.bin", "*.model"])
def get_hf_config(
model: str,
trust_remote_code: bool,
revision: Optional[str] = None,
model_override_args: Optional[dict] = None,
inference_args: Optional[dict] = None,
**kwargs,
):
is_gguf = check_gguf_file(model)
if is_gguf:
kwargs["gguf_file"] = model
model = Path(model).parent
config = AutoConfig.from_pretrained(
model, trust_remote_code=trust_remote_code, revision=revision, **kwargs
)
if config.model_type in _CONFIG_REGISTRY:
config_class = _CONFIG_REGISTRY[config.model_type]
config = config_class.from_pretrained(model, revision=revision)
# NOTE(HandH1998): Qwen2VL requires `_name_or_path` attribute in `config`.
setattr(config, "_name_or_path", model)
if model_override_args:
config.update(model_override_args)
# Special architecture mapping check for GGUF models
if is_gguf:
if config.model_type not in MODEL_FOR_CAUSAL_LM_MAPPING_NAMES:
raise RuntimeError(f"Can't get gguf config for {config.model_type}.")
model_type = MODEL_FOR_CAUSAL_LM_MAPPING_NAMES[config.model_type]
config.update({"architectures": [model_type]})
return config
def get_diffusers_config(
model: str,
inference_args: Optional[dict] = None,
) -> Dict[str, Any]:
"""Gets a configuration for the given diffusers model.
Args:
model: The model name or path.
inference_args: Optional inference arguments to override in the config.
Returns:
The loaded configuration.
"""
# Check if the model path exists
if os.path.exists(model):
config_file = os.path.join(model, "config.json")
if os.path.exists(config_file):
try:
# Load the config directly from the file
with open(config_file, "r") as f:
config_dict = json.load(f)
# TODO(will): apply any overrides from inference args
return config_dict
except Exception as e:
raise RuntimeError(f"Failed to load diffusers config from {config_file}: {e}")
else:
raise RuntimeError(f"Diffusers config file not found at {model}")
# Models don't use the same configuration key for determining the maximum
# context length. Store them here so we can sanely check them.
# NOTE: The ordering here is important. Some models have two of these and we
# have a preference for which value gets used.
CONTEXT_LENGTH_KEYS = [
"max_sequence_length",
"seq_length",
"max_seq_len",
"model_max_length",
"max_position_embeddings",
]
def get_context_length(config):
"""Get the context length of a model from a huggingface model configs."""
text_config = config
rope_scaling = getattr(text_config, "rope_scaling", None)
if rope_scaling:
rope_scaling_factor = rope_scaling.get("factor", 1)
if "original_max_position_embeddings" in rope_scaling:
rope_scaling_factor = 1
if rope_scaling.get("rope_type", None) == "llama3":
rope_scaling_factor = 1
else:
rope_scaling_factor = 1
for key in CONTEXT_LENGTH_KEYS:
val = getattr(text_config, key, None)
if val is not None:
return int(rope_scaling_factor * val)
return 2048
# A fast LLaMA tokenizer with the pre-processed `tokenizer.json` file.
_FAST_LLAMA_TOKENIZER = "hf-internal-testing/llama-tokenizer"
def get_tokenizer(
tokenizer_name: str,
*args,
tokenizer_mode: str = "auto",
trust_remote_code: bool = False,
tokenizer_revision: Optional[str] = None,
**kwargs,
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast]:
"""Gets a tokenizer for the given model name via Huggingface."""
if tokenizer_mode == "slow":
if kwargs.get("use_fast", False):
raise ValueError("Cannot use the fast tokenizer in slow tokenizer mode.")
kwargs["use_fast"] = False
is_gguf = check_gguf_file(tokenizer_name)
if is_gguf:
kwargs["gguf_file"] = tokenizer_name
tokenizer_name = Path(tokenizer_name).parent
try:
tokenizer = AutoTokenizer.from_pretrained(
tokenizer_name,
*args,
trust_remote_code=trust_remote_code,
tokenizer_revision=tokenizer_revision,
clean_up_tokenization_spaces=False,
**kwargs,
)
except TypeError as e:
# The LLaMA tokenizer causes a protobuf error in some environments.
err_msg = (
"Failed to load the tokenizer. If you are using a LLaMA V1 model "
f"consider using '{_FAST_LLAMA_TOKENIZER}' instead of the "
"original tokenizer."
)
raise RuntimeError(err_msg) from e
except ValueError as e:
# If the error pertains to the tokenizer class not existing or not
# currently being imported, suggest using the --trust-remote-code flag.
if not trust_remote_code and (
"does not exist or is not currently imported." in str(e)
or "requires you to execute the tokenizer file" in str(e)
):
err_msg = (
"Failed to load the tokenizer. If the tokenizer is a custom "
"tokenizer not yet available in the HuggingFace transformers "
"library, consider setting `trust_remote_code=True` in LLM "
"or using the `--trust-remote-code` flag in the CLI."
)
raise RuntimeError(err_msg) from e
else:
raise e
if not isinstance(tokenizer, PreTrainedTokenizerFast):
warnings.warn(
"Using a slow tokenizer. This might cause a significant "
"slowdown. Consider using a fast tokenizer instead."
)
attach_additional_stop_token_ids(tokenizer)
return tokenizer
def get_processor(
tokenizer_name: str,
*args,
tokenizer_mode: str = "auto",
trust_remote_code: bool = False,
tokenizer_revision: Optional[str] = None,
**kwargs,
):
processor = AutoProcessor.from_pretrained(
tokenizer_name,
*args,
trust_remote_code=trust_remote_code,
tokenizer_revision=tokenizer_revision,
**kwargs,
)
attach_additional_stop_token_ids(processor.tokenizer)
return processor
def attach_additional_stop_token_ids(tokenizer):
# Special handling for stop token <|eom_id|> generated by llama 3 tool use.
if "<|eom_id|>" in tokenizer.get_added_vocab():
tokenizer.additional_stop_token_ids = set(
[tokenizer.get_added_vocab()["<|eom_id|>"]]
)
else:
tokenizer.additional_stop_token_ids = None
def check_gguf_file(model: Union[str, os.PathLike]) -> bool:
"""Check if the file is a GGUF model."""
model = Path(model)
if not model.is_file():
return False
elif model.suffix == ".gguf":
return True
with open(model, "rb") as f:
header = f.read(4)
return header == b"GGUF"
@@ -1,409 +0,0 @@
from abc import ABC, abstractmethod
import dataclasses
import torch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.logger import init_logger
import os
import glob
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
from transformers import PretrainedConfig, AutoTokenizer
from fastvideo.v1.models.hf_transformer_utils import get_hf_config, get_diffusers_config
from fastvideo.v1.models import get_scheduler
from fastvideo.v1.models.registry import ModelRegistry
from safetensors.torch import load_file as safetensors_load_file
from typing import Tuple, List, Optional, Any, Generator
import time
import torch.nn as nn
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.v1.models.loader.weight_utils import (
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
pt_weights_iterator,
safetensors_weights_iterator)
from fastvideo.v1.models.loader.utils import set_default_torch_dtype
from typing import (Any, Dict, Generator, Iterable, List, Optional,
Tuple, cast)
logger = init_logger(__name__)
class ComponentLoader(ABC):
"""Base class for loading a specific type of model component."""
def __init__(self, device=None):
self.device = device
@abstractmethod
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
"""
Load the component based on the model path, architecture, and inference args.
Args:
model_path: Path to the component model
architecture: Architecture of the component model
inference_args: Inference arguments
Returns:
The loaded component
"""
raise NotImplementedError
@classmethod
def for_module_type(cls, module_type: str, transformers_or_diffusers: str) -> 'ComponentLoader':
"""
Factory method to create a component loader for a specific module type.
Args:
module_type: Type of module (e.g., "vae", "text_encoder", "transformer", "scheduler")
transformers_or_diffusers: Whether the module is from transformers or diffusers
Returns:
A component loader for the specified module type
"""
# Map of module types to their loader classes and expected library
module_loaders = {
"scheduler": (SchedulerLoader, "diffusers"),
"transformer": (TransformerLoader, "diffusers"),
"vae": (VAELoader, "diffusers"),
"text_encoder": (TextEncoderLoader, "transformers"),
"text_encoder_2": (TextEncoderLoader, "transformers"),
"tokenizer": (TokenizerLoader, "transformers"),
"tokenizer_2": (TokenizerLoader, "transformers"),
}
if module_type in module_loaders:
loader_cls, expected_library = module_loaders[module_type]
# Assert that the library matches what's expected for this module type
assert transformers_or_diffusers == expected_library, f"{module_type} must be loaded from {expected_library}, got {transformers_or_diffusers}"
return loader_cls()
# For unknown module types, use a generic loader
logger.warning(f"No specific loader found for module type: {module_type}. Using generic loader.")
return GenericComponentLoader(transformers_or_diffusers)
class TextEncoderLoader(ComponentLoader):
"""Loader for text encoders."""
@dataclasses.dataclass
class Source:
"""A source for weights."""
model_or_path: str
"""The model ID or path."""
prefix: str = ""
"""A prefix to prepend to all weights."""
fall_back_to_pt: bool = True
"""Whether .pt weights can be used."""
allow_patterns_overrides: Optional[list[str]] = None
"""If defined, weights will load exclusively using these patterns."""
counter_before_loading_weights: float = 0.0
counter_after_loading_weights: float = 0.0
def _prepare_weights(
self,
model_name_or_path: str,
fall_back_to_pt: bool,
allow_patterns_overrides: Optional[list[str]],
) -> Tuple[str, List[str], bool]:
"""Prepare weights for the model.
If the model is not local, it will be downloaded."""
# model_name_or_path = (self._maybe_download_from_modelscope(
# model_name_or_path, revision) or model_name_or_path)
is_local = os.path.isdir(model_name_or_path)
assert is_local, "Model path must be a local directory"
use_safetensors = False
index_file = SAFE_WEIGHTS_INDEX_NAME
allow_patterns = ["*.safetensors", "*.bin"]
if fall_back_to_pt:
allow_patterns += ["*.pt"]
if allow_patterns_overrides is not None:
allow_patterns = allow_patterns_overrides
hf_folder = model_name_or_path
hf_weights_files: List[str] = []
for pattern in allow_patterns:
hf_weights_files += glob.glob(os.path.join(hf_folder, pattern))
if len(hf_weights_files) > 0:
if pattern == "*.safetensors":
use_safetensors = True
break
if use_safetensors:
hf_weights_files = filter_duplicate_safetensors_files(
hf_weights_files, hf_folder, index_file)
else:
hf_weights_files = filter_files_not_needed_for_inference(
hf_weights_files)
if len(hf_weights_files) == 0:
raise RuntimeError(
f"Cannot find any model weights with `{model_name_or_path}`")
return hf_folder, hf_weights_files, use_safetensors
def _get_weights_iterator(
self, source: "Source"
) -> Generator[Tuple[str, torch.Tensor], None, None]:
"""Get an iterator for the model weights based on the load format."""
hf_folder, hf_weights_files, use_safetensors = self._prepare_weights(
source.model_or_path, source.fall_back_to_pt,
source.allow_patterns_overrides)
if use_safetensors:
weights_iterator = safetensors_weights_iterator(hf_weights_files)
else:
weights_iterator = pt_weights_iterator(hf_weights_files)
if self.counter_before_loading_weights == 0.0:
self.counter_before_loading_weights = time.perf_counter()
# Apply the prefix.
return ((source.prefix + name, tensor)
for (name, tensor) in weights_iterator)
def _get_all_weights(
self,
model_config: Dict[str, Any],
model: nn.Module,
) -> Generator[Tuple[str, torch.Tensor], None, None]:
primary_weights = TextEncoderLoader.Source(
model_config.model,
prefix="",
fall_back_to_pt=getattr(model, "fall_back_to_pt_during_load",
True),
allow_patterns_overrides=getattr(model, "allow_patterns_overrides",
None),
)
yield from self._get_weights_iterator(primary_weights)
secondary_weights = cast(
Iterable[TextEncoderLoader.Source],
getattr(model, "secondary_weights", ()),
)
for source in secondary_weights:
yield from self._get_weights_iterator(source)
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
"""Load the text encoders based on the model path, architecture, and inference args."""
model_config: PretrainedConfig = get_hf_config(
model=model_path,
trust_remote_code=inference_args.trust_remote_code,
revision=inference_args.revision,
model_override_args=None,
inference_args=inference_args,
)
logger.info(f"HF Model config: {model_config}")
target_device = torch.device(inference_args.device_str)
# TODO(will): add support for other dtypes
return self.load_model(model_path, model_config, target_device)
def load_model(self, model_path: str, model_config, target_device: torch.device):
with set_default_torch_dtype(torch.float16):
with target_device:
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
model = model_cls(model_config)
weights_to_load = {name for name, _ in model.named_parameters()}
model_config.model = model_path
loaded_weights = model.load_weights(
self._get_all_weights(model_config, model))
self.counter_after_loading_weights = time.perf_counter()
logger.info(
"Loading weights took %.2f seconds",
self.counter_after_loading_weights -
self.counter_before_loading_weights)
# We only enable strict check for non-quantized models
# that have loaded weights tracking currently.
# if loaded_weights is not None:
weights_not_loaded = weights_to_load - loaded_weights
if weights_not_loaded:
raise ValueError(
"Following weights were not initialized from "
f"checkpoint: {weights_not_loaded}")
# TODO(will): add support for training/finetune
return model.eval()
class TokenizerLoader(ComponentLoader):
"""Loader for tokenizers."""
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
"""Load the tokenizer based on the model path, architecture, and inference args."""
logger.info(f"Loading tokenizer from {model_path}")
tokenizer = AutoTokenizer.from_pretrained(
model_path,
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
# other method of config?
padding_size='right',
)
logger.info(f"Loaded tokenizer: {tokenizer.__class__.__name__}")
return tokenizer
class VAELoader(ComponentLoader):
"""Loader for VAE."""
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
"""Load the VAE based on the model path, architecture, and inference args."""
# TODO(will): move this to a constants file
from fastvideo.v1.utils import PRECISION_TO_TYPE
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name")
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
config.pop("_diffusers_version")
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(**config).to(inference_args.device)
# Find all safetensors files
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
# TODO(PY)
assert len(safetensors_list) == 1, f"Found {len(safetensors_list)} safetensors files in {d}"
loaded = safetensors_load_file(safetensors_list[0])
vae.load_state_dict(loaded)
dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
vae = vae.eval().to(dtype)
# TODO(will): should we define hunyuan vae config class?
vae_kwargs = {
"s_ratio": config["spatial_compression_ratio"],
"t_ratio": config["temporal_compression_ratio"],
}
vae.kwargs = vae_kwargs
return vae
class TransformerLoader(ComponentLoader):
"""Loader for transformer."""
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
"""Load the transformer based on the model path, architecture, and inference args."""
model_config = get_diffusers_config(model=model_path)
cls_name = model_config.pop("_class_name")
if cls_name is None:
raise ValueError(f"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
model_config.pop("_diffusers_version")
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
# Find all safetensors files
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
logger.info(f"Loading model from {len(safetensors_list)} safetensors files in {model_path}")
# initialize_sequence_parallel_group(inference_args.sp_size)
# Load the model using FSDP loader
logger.info(f"Loading model from {cls_name}")
model = load_fsdp_model(
model_cls=model_cls,
init_params=model_config,
weight_dir_list=safetensors_list,
device=inference_args.device,
cpu_offload=inference_args.use_cpu_offload
)
total_params = sum(p.numel() for p in model.parameters())
logger.info(f"Loaded model with {total_params / 1e9:.2f}B parameters")
model.eval()
return model
class SchedulerLoader(ComponentLoader):
"""Loader for scheduler."""
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
"""Load the scheduler based on the model path, architecture, and inference args."""
scheduler = get_scheduler(
module_path=model_path,
architecture=architecture,
inference_args=inference_args,
)
logger.info(f"Scheduler loaded: {scheduler}")
return scheduler
class GenericComponentLoader(ComponentLoader):
"""Generic loader for components that don't have a specific loader."""
def __init__(self, library="transformers"):
super().__init__()
self.library = library
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
"""Load a generic component based on the model path, architecture, and inference args."""
logger.warning(f"Using generic loader for {model_path} with library {self.library}")
if self.library == "transformers":
from transformers import AutoModel
model = AutoModel.from_pretrained(
model_path,
trust_remote_code=inference_args.trust_remote_code,
revision=inference_args.revision,
)
logger.info(f"Loaded generic transformers model: {model.__class__.__name__}")
return model
elif self.library == "diffusers":
logger.warning(f"Generic loading for diffusers components is not fully implemented")
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
model_config = get_diffusers_config(model=model_path)
logger.info(f"Diffusers Model config: {model_config}")
# This is a placeholder - in a real implementation, you'd need to handle this properly
return None
else:
raise ValueError(f"Unsupported library: {self.library}")
class PipelineComponentLoader:
"""
Utility class for loading pipeline components.
This replaces the chain of if-else statements in load_pipeline_module.
"""
@staticmethod
def load_module(module_name: str, component_model_path: str, transformers_or_diffusers: str,
architecture: str, inference_args: InferenceArgs):
"""
Load a pipeline module.
Args:
module_name: Name of the module (e.g., "vae", "text_encoder", "transformer", "scheduler")
component_model_path: Path to the component model
transformers_or_diffusers: Whether the module is from transformers or diffusers
architecture: Architecture of the component model
inference_args: Inference arguments
Returns:
The loaded module
"""
logger.info(f"Loading {module_name} using {transformers_or_diffusers} from {component_model_path}")
# Get the appropriate loader for this module type
loader = ComponentLoader.for_module_type(module_name, transformers_or_diffusers)
# Load the module
return loader.load(component_model_path, architecture, inference_args)
-227
View File
@@ -1,227 +0,0 @@
# Adapted from torchtune
# Copyright 2024 The TorchTune Authors.
# Copyright 2025 The FastVideo Authors.
from collections import defaultdict
from typing import Any, Callable, Dict, Generator, List, Optional, Tuple, Type
from itertools import chain
import torch
from torch import nn
from torch.distributed import DeviceMesh, init_device_mesh
from fastvideo.v1.distributed.parallel_state import get_sequence_model_parallel_world_size
from torch.distributed._composable.fsdp import CPUOffloadPolicy, fully_shard
from torch.distributed._tensor import distribute_tensor
from torch.nn.modules.module import _IncompatibleKeys
from vllm.model_executor.model_loader.weight_utils import safetensors_weights_iterator
import contextlib
import re
# TODO(PY): move this to utils elsewhere
@contextlib.contextmanager
def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
"""
Context manager to set torch's default dtype.
Args:
dtype (torch.dtype): The desired default dtype inside the context manager.
Returns:
ContextManager: context manager for setting default dtype.
Example:
>>> with set_default_dtype(torch.bfloat16):
>>> x = torch.tensor([1, 2, 3])
>>> x.dtype
torch.bfloat16
"""
old_dtype = torch.get_default_dtype()
torch.set_default_dtype(dtype)
try:
yield
finally:
torch.set_default_dtype(old_dtype)
def get_param_names_mapping(mapping_dict: Dict[str, str]) -> Callable[[str], str]:
"""
Creates a mapping function that transforms parameter names using regex patterns.
Args:
mapping_dict (Dict[str, str]): Dictionary mapping regex patterns to replacement patterns
param_name (str): The parameter name to be transformed
Returns:
Callable[[str], str]: A function that maps parameter names from source to target format
"""
def mapping_fn(name: str) -> str:
# Try to match and transform the name using the regex patterns in mapping_dict
for pattern, replacement in mapping_dict.items():
match = re.match(pattern, name)
if match:
merge_index = None
total_splitted_params = None
if isinstance(replacement, tuple):
merge_index = replacement[1]
total_splitted_params = replacement[2]
replacement= replacement[0]
name = re.sub(pattern, replacement, name)
return name, merge_index, total_splitted_params
# If no pattern matches, return the original name
return name, None, None
return mapping_fn
# TODO(PY): add compile option
def load_fsdp_model(
model_cls: Type[nn.Module],
init_params: Dict[str, Any],
weight_dir_list: List[str],
device: torch.device,
cpu_offload: bool = False,
default_dtype: Optional[torch.dtype] = torch.bfloat16,
) -> torch.nn.Module:
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
device_mesh = init_device_mesh(
"cuda",
mesh_shape=(get_sequence_model_parallel_world_size(),),
mesh_dim_names=("dp", ),
)
shard_model(model, cpu_offload=cpu_offload, reshard_after_forward=True, dp_mesh=device_mesh["dp"])
weight_iterator = safetensors_weights_iterator(weight_dir_list)
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
load_fsdp_model_from_full_model_state_dict(
model,
weight_iterator,
device,
strict=True,
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
)
for n, p in chain(model.named_parameters(), model.named_buffers()):
if p.is_meta:
raise RuntimeError(f"Unexpected param or buffer {n} on meta device.")
for p in model.parameters():
p.requires_grad = False
return model
def shard_model(
model,
*,
cpu_offload: bool,
reshard_after_forward: bool = True,
dp_mesh: Optional[DeviceMesh] = None,
) -> None:
"""
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API.
This method will over the model's named modules from the bottom-up and apply shard modules
based on whether they meet any of the criteria from shard_conditions.
Args:
model (TransformerDecoder): Model to shard with FSDP.
shard_conditions (List[Callable[[str, nn.Module], bool]]): A list of functions to determine
which modules to shard with FSDP. Each function should take module name (relative to root)
and the module itself, returning True if FSDP should shard the module and False otherwise.
If any of shard_conditions return True for a given module, it will be sharded by FSDP.
cpu_offload (bool): If set to True, FSDP will offload parameters, gradients, and optimizer
states to CPU.
reshard_after_forward (bool): Whether to reshard parameters and buffers after
the forward pass. Setting this to True corresponds to the FULL_SHARD sharding strategy
from FSDP1, while setting it to False corresponds to the SHARD_GRAD_OP sharding strategy.
dp_mesh (Optional[DeviceMesh]): Device mesh to use for FSDP sharding under mutliple parallelism.
Default to None.
Raises:
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
"""
fsdp_kwargs = {"reshard_after_forward": reshard_after_forward, "mesh": dp_mesh}
if cpu_offload:
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
# Shard the model with FSDP, iterating in reverse to start with
# lowest-level modules first
num_layers_sharded = 0
for n, m in reversed(list(model.named_modules())):
if any([shard_condition(n, m) for shard_condition in model._fsdp_shard_conditions]):
fully_shard(m, **fsdp_kwargs)
num_layers_sharded += 1
if num_layers_sharded == 0:
raise ValueError(
"No layer modules were sharded. Please check if shard conditions are working as expected."
)
# Finally shard the entire model to account for any stragglers
fully_shard(model, **fsdp_kwargs)
# TODO(PY): device mesh for cfg parallel
def load_fsdp_model_from_full_model_state_dict(
model: torch.nn.Module,
full_sd_iterator: Generator[Tuple[str, torch.Tensor], None, None],
device: torch.device,
strict: bool = False,
cpu_offload: bool = False,
param_names_mapping: Optional[Callable[[str], str]] = None,
) -> _IncompatibleKeys:
"""
Converting full state dict into a sharded state dict
and loading it into FSDP model
Args:
model (FSDPModule): Model to generate fully qualified names for cpu_state_dict
full_sd_iterator (Generator): an iterator yielding (param_name, tensor) pairs
device (torch.device): device used to move full state dict tensors
strict (bool): flag to check if to load the model in strict mode
cpu_offload (bool): flag to check if offload to CPU is enabled
param_names_mapping (Optional[Callable[[str], str]]): a function that maps full param name to sharded param name
Returns:
``NamedTuple`` with ``missing_keys`` and ``unexpected_keys`` fields:
* **missing_keys** is a list of str containing the missing keys
* **unexpected_keys** is a list of str containing the unexpected keys
Raises:
NotImplementedError: If got FSDP with more than 1D.
"""
meta_sharded_sd = model.state_dict()
sharded_sd = {}
to_merge_params = defaultdict(dict)
for source_param_name, full_tensor in full_sd_iterator:
target_param_name, merge_index, num_params_to_merge = param_names_mapping(source_param_name)
if merge_index is not None:
to_merge_params[target_param_name][merge_index] = full_tensor
if len(to_merge_params[target_param_name]) == num_params_to_merge:
# cat at dim=1 according to the merge_index order
sorted_tensors = [to_merge_params[target_param_name][i] for i in range(num_params_to_merge)]
full_tensor = torch.cat(sorted_tensors, dim=0)
del to_merge_params[target_param_name]
else:
continue
sharded_meta_param = meta_sharded_sd.get(target_param_name)
if sharded_meta_param is None:
raise ValueError(f"Parameter {source_param_name}-->{target_param_name} not found in meta sharded state dict")
full_tensor = full_tensor.to(sharded_meta_param.dtype).to(device)
if not hasattr(sharded_meta_param, "device_mesh"):
# In cases where parts of the model aren't sharded, some parameters will be plain tensors
sharded_tensor = full_tensor
else:
sharded_tensor = distribute_tensor(
full_tensor,
sharded_meta_param.device_mesh,
sharded_meta_param.placements,
)
if cpu_offload:
sharded_tensor = sharded_tensor.cpu()
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
# choose `assign=True` since we cannot call `copy_` on meta tensor
return model.load_state_dict(sharded_sd, strict=strict, assign=True)
-17
View File
@@ -1,17 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Utilities for selecting and loading models."""
import contextlib
import torch
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@contextlib.contextmanager
def set_default_torch_dtype(dtype: torch.dtype):
"""Sets the default torch dtype to the given dtype."""
old_dtype = torch.get_default_dtype()
torch.set_default_dtype(dtype)
yield
torch.set_default_dtype(old_dtype)
-346
View File
@@ -1,346 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm
# Copyright 2023 The vLLM Authors.
# Copyright 2025 The FastVideo Authors.
"""Utilities for downloading and initializing model weights."""
import fnmatch
import hashlib
import json
import os
import tempfile
import time
from collections import defaultdict
from pathlib import Path
from typing import Generator, List, Optional, Tuple, Union
import filelock
import huggingface_hub.constants
import torch
from huggingface_hub import HfFileSystem, hf_hub_download, snapshot_download
from safetensors.torch import safe_open
from tqdm.auto import tqdm
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
# use system-level temp directory for file locks, so that multiple users
# can share the same lock without error.
# lock files in the temp directory will be automatically deleted when the
# system reboots, so users will not complain about annoying lock files
temp_dir = tempfile.gettempdir()
def enable_hf_transfer():
"""automatically activates hf_transfer
"""
if "HF_HUB_ENABLE_HF_TRANSFER" not in os.environ:
try:
# enable hf hub transfer if available
import hf_transfer # type: ignore # noqa
huggingface_hub.constants.HF_HUB_ENABLE_HF_TRANSFER = True
except ImportError:
pass
enable_hf_transfer()
class DisabledTqdm(tqdm):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs, disable=True)
def get_lock(model_name_or_path: Union[str, Path],
cache_dir: Optional[str] = None):
lock_dir = cache_dir or temp_dir
model_name_or_path = str(model_name_or_path)
os.makedirs(os.path.dirname(lock_dir), exist_ok=True)
model_name = model_name_or_path.replace("/", "-")
hash_name = hashlib.sha256(model_name.encode()).hexdigest()
# add hash to avoid conflict with old users' lock files
lock_file_name = hash_name + model_name + ".lock"
# mode 0o666 is required for the filelock to be shared across users
lock = filelock.FileLock(os.path.join(lock_dir, lock_file_name),
mode=0o666)
return lock
def _shared_pointers(tensors):
ptrs = defaultdict(list)
for k, v in tensors.items():
ptrs[v.data_ptr()].append(k)
failing = []
for _, names in ptrs.items():
if len(names) > 1:
failing.append(names)
return failing
def download_weights_from_hf(
model_name_or_path: str,
cache_dir: Optional[str],
allow_patterns: List[str],
revision: Optional[str] = None,
ignore_patterns: Optional[Union[str, List[str]]] = None,
) -> str:
"""Download model weights from Hugging Face Hub.
Args:
model_name_or_path (str): The model name or path.
cache_dir (Optional[str]): The cache directory to store the model
weights. If None, will use HF defaults.
allow_patterns (List[str]): The allowed patterns for the
weight files. Files matched by any of the patterns will be
downloaded.
revision (Optional[str]): The revision of the model.
ignore_patterns (Optional[Union[str, List[str]]]): The patterns to
filter out the weight files. Files matched by any of the patterns
will be ignored.
Returns:
str: The path to the downloaded model weights.
"""
local_only = huggingface_hub.constants.HF_HUB_OFFLINE
if not local_only:
# Before we download we look at that is available:
fs = HfFileSystem()
file_list = fs.ls(model_name_or_path, detail=False, revision=revision)
# depending on what is available we download different things
for pattern in allow_patterns:
matching = fnmatch.filter(file_list, pattern)
if len(matching) > 0:
allow_patterns = [pattern]
break
logger.info("Using model weights format %s", allow_patterns)
# Use file lock to prevent multiple processes from
# downloading the same model weights at the same time.
with get_lock(model_name_or_path, cache_dir):
start_time = time.perf_counter()
hf_folder = snapshot_download(
model_name_or_path,
allow_patterns=allow_patterns,
ignore_patterns=ignore_patterns,
cache_dir=cache_dir,
tqdm_class=DisabledTqdm,
revision=revision,
local_files_only=local_only,
)
time_taken = time.perf_counter() - start_time
if time_taken > 0.5:
logger.info("Time spent downloading weights for %s: %.6f seconds",
model_name_or_path, time_taken)
return hf_folder
def download_safetensors_index_file_from_hf(
model_name_or_path: str,
index_file: str,
cache_dir: Optional[str],
revision: Optional[str] = None,
) -> None:
"""Download hf safetensors index file from Hugging Face Hub.
Args:
model_name_or_path (str): The model name or path.
cache_dir (Optional[str]): The cache directory to store the model
weights. If None, will use HF defaults.
revision (Optional[str]): The revision of the model.
"""
# Use file lock to prevent multiple processes from
# downloading the same model weights at the same time.
with get_lock(model_name_or_path, cache_dir):
try:
# Download the safetensors index file.
hf_hub_download(
repo_id=model_name_or_path,
filename=index_file,
cache_dir=cache_dir,
revision=revision,
local_files_only=huggingface_hub.constants.HF_HUB_OFFLINE,
)
# If file not found on remote or locally, we should not fail since
# only some models will have index_file.
except huggingface_hub.utils.EntryNotFoundError:
logger.info("No %s found in remote.", index_file)
except huggingface_hub.utils.LocalEntryNotFoundError:
logger.info("No %s found in local cache.", index_file)
# For models like Mistral-7B-v0.3, there are both sharded
# safetensors files and a consolidated safetensors file.
# Passing both of these to the weight loader functionality breaks.
# So, we use the index_file to
# look up which safetensors files should be used.
def filter_duplicate_safetensors_files(hf_weights_files: List[str],
hf_folder: str,
index_file: str) -> List[str]:
# model.safetensors.index.json is a mapping from keys in the
# torch state_dict to safetensors file holding that weight.
index_file_name = os.path.join(hf_folder, index_file)
if not os.path.isfile(index_file_name):
return hf_weights_files
# Iterate through the weight_map (weight_name: safetensors files)
# to identify weights that we should use.
with open(index_file_name) as f:
weight_map = json.load(f)["weight_map"]
weight_files_in_index = set()
for weight_name in weight_map:
weight_files_in_index.add(
os.path.join(hf_folder, weight_map[weight_name]))
# Filter out any fields that are not found in the index file.
hf_weights_files = [
f for f in hf_weights_files if f in weight_files_in_index
]
return hf_weights_files
def filter_files_not_needed_for_inference(
hf_weights_files: List[str]) -> List[str]:
"""
Exclude files that are not needed for inference.
See https://github.com/huggingface/transformers/blob/v4.34.0/src/transformers/trainer.py#L227-L233
"""
blacklist = [
"training_args.bin",
"optimizer.bin",
"optimizer.pt",
"scheduler.pt",
"scaler.pt",
]
hf_weights_files = [
f for f in hf_weights_files
if not any(f.endswith(x) for x in blacklist)
]
return hf_weights_files
# explicitly use pure text format, with a newline at the end
# this makes it impossible to see the animation in the progress bar
# but will avoid messing up with ray or multiprocessing, which wraps
# each line of output with some prefix.
_BAR_FORMAT = "{desc}: {percentage:3.0f}% Completed | {n_fmt}/{total_fmt} [{elapsed}<{remaining}, {rate_fmt}]\n" # noqa: E501
def safetensors_weights_iterator(
hf_weights_files: List[str]
) -> Generator[Tuple[str, torch.Tensor], None, None]:
"""Iterate over the weights in the model safetensor files."""
enable_tqdm = not torch.distributed.is_initialized(
) or torch.distributed.get_rank() == 0
for st_file in tqdm(
hf_weights_files,
desc="Loading safetensors checkpoint shards",
disable=not enable_tqdm,
bar_format=_BAR_FORMAT,
):
with safe_open(st_file, framework="pt") as f:
for name in f.keys(): # noqa: SIM118
param = f.get_tensor(name)
yield name, param
def pt_weights_iterator(
hf_weights_files: List[str]
) -> Generator[Tuple[str, torch.Tensor], None, None]:
"""Iterate over the weights in the model bin/pt files."""
enable_tqdm = not torch.distributed.is_initialized(
) or torch.distributed.get_rank() == 0
for bin_file in tqdm(
hf_weights_files,
desc="Loading pt checkpoint shards",
disable=not enable_tqdm,
bar_format=_BAR_FORMAT,
):
state = torch.load(bin_file, map_location="cpu", weights_only=True)
yield from state.items()
del state
def default_weight_loader(param: torch.Tensor,
loaded_weight: torch.Tensor) -> None:
"""Default weight loader."""
try:
if param.numel() == 1 and loaded_weight.numel() == 1:
# Sometimes scalar values aren't considered tensors with shapes
# so if both param and loaded_weight are a scalar,
# "broadcast" instead of copy
param.data.fill_(loaded_weight.item())
else:
assert param.size() == loaded_weight.size(), (
f"Attempted to load weight ({loaded_weight.size()}) "
f"into parameter ({param.size()})")
param.data.copy_(loaded_weight)
except Exception:
# NOTE: This exception is added for the purpose of setting breakpoint to
# debug weight loading issues.
raise
def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> Optional[str]:
"""Remap the name of FP8 k/v_scale parameters.
This function handles the remapping of FP8 k/v_scale parameter names.
It detects if the given name ends with a suffix and attempts to remap
it to the expected name format in the model. If the remapped name is not
found in the params_dict, a warning is printed and None is returned.
Args:
name (str): The original loaded checkpoint parameter name.
params_dict (dict): Dictionary containing the model's named parameters.
Returns:
str: The remapped parameter name if successful, or the original name
if no remapping is needed.
None: If the remapped name is not found in params_dict.
"""
if name.endswith(".kv_scale"):
logger.warning_once(
"DEPRECATED. Found kv_scale in the checkpoint. "
"This format is deprecated in favor of separate k_scale and "
"v_scale tensors and will be removed in a future release. "
"Functionally, we will remap kv_scale to k_scale and duplicate "
"k_scale to v_scale")
# NOTE: we remap the deprecated kv_scale to k_scale
remapped_name = name.replace(".kv_scale", ".attn.k_scale")
if remapped_name not in params_dict:
logger.warning_once(
f"Found kv_scale in the checkpoint (e.g. {name}), "
"but not found the expected name in the model "
f"(e.g. {remapped_name}). kv_scale is "
"not loaded.")
return None
return remapped_name
possible_scale_names = [".k_scale", ".v_scale"]
modelopt_scale_names = [
".self_attn.k_proj.k_scale", ".self_attn.v_proj.v_scale"
]
for scale_name in possible_scale_names:
if name.endswith(scale_name):
if any(mo_scale_name in name
for mo_scale_name in modelopt_scale_names):
remapped_name = name.replace(
f".self_attn.{scale_name[1]}_proj{scale_name}",
f".self_attn.attn{scale_name}")
else:
remapped_name = name.replace(scale_name, f".attn{scale_name}")
if remapped_name not in params_dict:
logger.warning_once(
f"Found {scale_name} in the checkpoint (e.g. {name}), "
"but not found the expected name in the model "
f"(e.g. {remapped_name}). {scale_name} is "
"not loaded.")
return None
return remapped_name
# If there were no matches, return the untouched param name
return name
-434
View File
@@ -1,434 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/parameter.py
from fractions import Fraction
from typing import Callable, Optional, Union
import torch
from torch.nn import Parameter
from fastvideo.v1.distributed import get_tensor_model_parallel_rank
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.utils import _make_synced_weight_loader
__all__ = [
"BasevLLMParameter", "PackedvLLMParameter", "PerTensorScaleParameter",
"ModelWeightParameter", "ChannelQuantScaleParameter",
"GroupQuantScaleParameter", "PackedColumnParameter", "RowvLLMParameter"
]
logger = init_logger(__name__)
class BasevLLMParameter(Parameter):
"""
Base parameter for vLLM linear layers. Extends the torch.nn.parameter
by taking in a linear weight loader. Will copy the loaded weight
into the parameter when the provided weight loader is called.
"""
def __new__(cls, data: torch.Tensor, **kwargs):
return super().__new__(cls, data=data, requires_grad=False)
def __init__(self, data: torch.Tensor, weight_loader: Callable):
"""
Initialize the BasevLLMParameter
:param data: torch tensor with the parameter data
:param weight_loader: weight loader callable
:returns: a torch.nn.parameter
"""
# During weight loading, we often do something like:
# narrowed_tensor = param.data.narrow(0, offset, len)
# narrowed_tensor.copy_(real_weight)
# expecting narrowed_tensor and param.data to share the same storage.
# However, on TPUs, narrowed_tensor will lazily propagate to the base
# tensor, which is param.data, leading to the redundant memory usage.
# This sometimes causes OOM errors during model loading. To avoid this,
# we sync the param tensor after its weight loader is called.
from vllm.platforms import current_platform
if current_platform.is_tpu():
weight_loader = _make_synced_weight_loader(weight_loader)
self._weight_loader = weight_loader
@property
def weight_loader(self):
return self._weight_loader
def _is_1d_and_scalar(self, loaded_weight: torch.Tensor):
cond1 = self.data.ndim == 1 and self.data.numel() == 1
cond2 = loaded_weight.ndim == 0 and loaded_weight.numel() == 1
return (cond1 and cond2)
def _assert_and_load(self, loaded_weight: torch.Tensor):
assert (self.data.shape == loaded_weight.shape
or self._is_1d_and_scalar(loaded_weight))
self.data.copy_(loaded_weight)
def load_column_parallel_weight(self, loaded_weight: torch.Tensor):
self._assert_and_load(loaded_weight)
def load_row_parallel_weight(self, loaded_weight: torch.Tensor):
self._assert_and_load(loaded_weight)
def load_merged_column_weight(self, loaded_weight: torch.Tensor, **kwargs):
self._assert_and_load(loaded_weight)
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs):
self._assert_and_load(loaded_weight)
class _ColumnvLLMParameter(BasevLLMParameter):
"""
Private class defining weight loading functionality
(load_merged_column_weight, load_qkv_weight)
for parameters being loaded into linear layers with column
parallelism. This includes QKV and MLP layers which are
not already fused on disk. Requires an output dimension
to be defined. Called within the weight loader of
each of the column parallel linear layers.
"""
def __init__(self, output_dim: int, **kwargs):
self._output_dim = output_dim
super().__init__(**kwargs)
@property
def output_dim(self):
return self._output_dim
def load_column_parallel_weight(self, loaded_weight: torch.Tensor):
tp_rank = get_tensor_model_parallel_rank()
shard_size = self.data.shape[self.output_dim]
loaded_weight = loaded_weight.narrow(self.output_dim,
tp_rank * shard_size, shard_size)
assert self.data.shape == loaded_weight.shape
self.data.copy_(loaded_weight)
def load_merged_column_weight(self, loaded_weight: torch.Tensor, **kwargs):
shard_offset = kwargs.get("shard_offset")
shard_size = kwargs.get("shard_size")
if isinstance(
self,
(PackedColumnParameter,
PackedvLLMParameter)) and self.packed_dim == self.output_dim:
shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
shard_offset=shard_offset, shard_size=shard_size)
param_data = self.data
tp_rank = get_tensor_model_parallel_rank()
param_data = param_data.narrow(self.output_dim, shard_offset,
shard_size)
loaded_weight = loaded_weight.narrow(self.output_dim,
tp_rank * shard_size, shard_size)
assert param_data.shape == loaded_weight.shape
param_data.copy_(loaded_weight)
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs):
shard_offset = kwargs.get("shard_offset")
shard_size = kwargs.get("shard_size")
shard_id = kwargs.get("shard_id")
num_heads = kwargs.get("num_heads")
if isinstance(
self,
(PackedColumnParameter,
PackedvLLMParameter)) and self.output_dim == self.packed_dim:
shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
shard_offset=shard_offset, shard_size=shard_size)
param_data = self.data
tp_rank = get_tensor_model_parallel_rank()
shard_id = tp_rank if shard_id == "q" else tp_rank // num_heads
param_data = param_data.narrow(self.output_dim, shard_offset,
shard_size)
loaded_weight = loaded_weight.narrow(self.output_dim,
shard_id * shard_size, shard_size)
assert param_data.shape == loaded_weight.shape
param_data.copy_(loaded_weight)
class RowvLLMParameter(BasevLLMParameter):
"""
Parameter class defining weight_loading functionality
(load_row_parallel_weight) for parameters being loaded
into linear layers with row parallel functionality.
Requires an input_dim to be defined.
"""
def __init__(self, input_dim: int, **kwargs):
self._input_dim = input_dim
super().__init__(**kwargs)
@property
def input_dim(self):
return self._input_dim
def load_row_parallel_weight(self, loaded_weight: torch.Tensor):
tp_rank = get_tensor_model_parallel_rank()
shard_size = self.data.shape[self.input_dim]
loaded_weight = loaded_weight.narrow(self.input_dim,
tp_rank * shard_size, shard_size)
if len(loaded_weight.shape) == 0:
loaded_weight = loaded_weight.reshape(1)
assert self.data.shape == loaded_weight.shape
self.data.copy_(loaded_weight)
class ModelWeightParameter(_ColumnvLLMParameter, RowvLLMParameter):
"""
Parameter class for linear layer weights. Uses both column and
row parallelism.
"""
pass
class GroupQuantScaleParameter(_ColumnvLLMParameter, RowvLLMParameter):
"""
Parameter class for weight scales loaded for weights with
grouped quantization. Uses both column and row parallelism.
"""
pass
class ChannelQuantScaleParameter(_ColumnvLLMParameter):
"""
Parameter class for weight scales loaded for weights with
channel-wise quantization. Equivalent to _ColumnvLLMParameter.
"""
pass
class PerTensorScaleParameter(BasevLLMParameter):
"""
Parameter class for scales where the number of scales is
equivalent to the number of logical matrices in fused linear
layers (e.g. for QKV, there are 3 scales loaded from disk).
This is relevant to weights with per-tensor quantization.
Adds functionality to map the scalers to a shard during
weight loading.
Note: additional parameter manipulation may be handled
for each quantization config specifically, within
process_weights_after_loading
"""
def __init__(self, **kwargs):
self.qkv_idxs = {"q": 0, "k": 1, "v": 2}
super().__init__(**kwargs)
def _shard_id_as_int(self, shard_id: Union[str, int]) -> int:
if isinstance(shard_id, int):
return shard_id
# if not int, assume shard_id for qkv
# map to int and return
assert isinstance(shard_id, str)
assert shard_id in self.qkv_idxs
return self.qkv_idxs[shard_id]
# For row parallel layers, no sharding needed
# load weight into parameter as is
def load_row_parallel_weight(self, *args, **kwargs):
super().load_row_parallel_weight(*args, **kwargs)
def load_merged_column_weight(self, *args, **kwargs):
self._load_into_shard_id(*args, **kwargs)
def load_qkv_weight(self, *args, **kwargs):
self._load_into_shard_id(*args, **kwargs)
def load_column_parallel_weight(self, *args, **kwargs):
super().load_row_parallel_weight(*args, **kwargs)
def _load_into_shard_id(self, loaded_weight: torch.Tensor,
shard_id: Union[str, int], **kwargs):
"""
Slice the parameter data based on the shard id for
loading.
"""
param_data = self.data
shard_id = self._shard_id_as_int(shard_id)
# AutoFP8 scales do not have a shape
# compressed-tensors scales do have a shape
if len(loaded_weight.shape) != 0:
assert loaded_weight.shape[0] == 1
loaded_weight = loaded_weight[0]
param_data = param_data[shard_id]
assert param_data.shape == loaded_weight.shape
param_data.copy_(loaded_weight)
class PackedColumnParameter(_ColumnvLLMParameter):
"""
Parameter for model parameters which are packed on disk
and support column parallelism only. See PackedvLLMParameter
for more details on the packed properties.
"""
def __init__(self,
packed_factor: Union[int, Fraction],
packed_dim: int,
marlin_tile_size: Optional[int] = None,
**kwargs):
self._packed_factor = packed_factor
self._packed_dim = packed_dim
self._marlin_tile_size = marlin_tile_size
super().__init__(**kwargs)
@property
def packed_dim(self):
return self._packed_dim
@property
def packed_factor(self):
return self._packed_factor
@property
def marlin_tile_size(self):
return self._marlin_tile_size
def adjust_shard_indexes_for_packing(self, shard_size, shard_offset):
return _adjust_shard_indexes_for_packing(
shard_size=shard_size,
shard_offset=shard_offset,
packed_factor=self.packed_factor,
marlin_tile_size=self.marlin_tile_size)
class PackedvLLMParameter(ModelWeightParameter):
"""
Parameter for model weights which are packed on disk.
Example: GPTQ Marlin weights are int4 or int8, packed into int32.
Extends the ModelWeightParameter to take in the
packed factor, the packed dimension, and optionally, marlin
tile size for marlin kernels. Adjusts the shard_size and
shard_offset for fused linear layers model weight loading
by accounting for packing and optionally, marlin tile size.
"""
def __init__(self,
packed_factor: Union[int, Fraction],
packed_dim: int,
marlin_tile_size: Optional[int] = None,
**kwargs):
self._packed_factor = packed_factor
self._packed_dim = packed_dim
self._marlin_tile_size = marlin_tile_size
super().__init__(**kwargs)
@property
def packed_dim(self):
return self._packed_dim
@property
def packed_factor(self):
return self._packed_factor
@property
def marlin_tile_size(self):
return self._marlin_tile_size
def adjust_shard_indexes_for_packing(self, shard_size, shard_offset):
return _adjust_shard_indexes_for_packing(
shard_size=shard_size,
shard_offset=shard_offset,
packed_factor=self.packed_factor,
marlin_tile_size=self.marlin_tile_size)
class BlockQuantScaleParameter(_ColumnvLLMParameter, RowvLLMParameter):
"""
Parameter class for weight scales loaded for weights with
block-wise quantization. Uses both column and row parallelism.
"""
pass
def permute_param_layout_(param: BasevLLMParameter, input_dim: int,
output_dim: int, **kwargs) -> BasevLLMParameter:
"""
Permute a parameter's layout to the specified input and output dimensions,
useful for forcing the parameter into a known layout, for example, if I need
a packed (quantized) weight matrix to be in the layout
{input_dim = 0, output_dim = 1, packed_dim = 0}
then I can call:
permute_param_layout_(x, input_dim=0, output_dim=1, packed_dim=0)
to ensure x is in the correct layout (permuting it to the correct layout if
required, asserting if it cannot get it to the correct layout)
"""
curr_input_dim = getattr(param, "input_dim", None)
curr_output_dim = getattr(param, "output_dim", None)
if curr_input_dim is None or curr_output_dim is None:
assert param.data.dim() == 2,\
"permute_param_layout_ only supports 2D parameters when either "\
"input_dim or output_dim is not set"
# if one of the dimensions is not set, set it to the opposite of the other
# we can only do this since we asserted the parameter is 2D above
if curr_input_dim is None:
assert curr_output_dim is not None,\
"either input or output dim must be set"
curr_input_dim = (curr_output_dim + 1) % 2
if curr_output_dim is None:
assert curr_input_dim is not None,\
"either input or output dim must be set"
curr_output_dim = (curr_input_dim + 1) % 2
# create permutation from the current layout to the layout with
# self.input_dim at input_dim and self.output_dim at output_dim preserving
# other dimensions
perm = [
i for i in range(param.data.dim())
if i not in [curr_input_dim, curr_output_dim]
]
perm.insert(input_dim, curr_input_dim)
perm.insert(output_dim, curr_output_dim)
if "packed_dim" in kwargs:
assert hasattr(param, "packed_dim") and\
param.packed_dim == perm[kwargs["packed_dim"]],\
"permute_param_layout_ currently doesn't support repacking"
param.data = param.data.permute(*perm)
if hasattr(param, "_input_dim"):
param._input_dim = input_dim
if hasattr(param, "_output_dim"):
param._output_dim = output_dim
if "packed_dim" in kwargs and hasattr(param, "_packed_dim"):
param._packed_dim = kwargs["packed_dim"]
return param
def _adjust_shard_indexes_for_marlin(shard_size, shard_offset,
marlin_tile_size):
return shard_size * marlin_tile_size, shard_offset * marlin_tile_size
def _adjust_shard_indexes_for_packing(shard_size, shard_offset, packed_factor,
marlin_tile_size):
shard_size = shard_size // packed_factor
shard_offset = shard_offset // packed_factor
if marlin_tile_size is not None:
return _adjust_shard_indexes_for_marlin(
shard_size=shard_size,
shard_offset=shard_offset,
marlin_tile_size=marlin_tile_size)
return shard_size, shard_offset
-293
View File
@@ -1,293 +0,0 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
import os
import pickle
import subprocess
import sys
import tempfile
from typing import AbstractSet, Callable, Dict, List, Optional, Tuple, Type, Union, TypeVar
import importlib
from functools import lru_cache
import cloudpickle
from torch import nn
from fastvideo.v1.logger import logger
# huggingface class name: (component_name, fastvideo module name, fastvideo class name)
_TEXT_TO_VIDEO_DIT_MODELS = {
"HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
}
_TEXT_ENCODER_MODELS = {
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
"LlamaModel": ("encoders", "llama", "LlamaModel"),
}
_IMAGE_ENCODER_MODELS = {
# "HunyuanVideoTransformer3DModel": ("image_encoder", "hunyuanvideo", "HunyuanVideoImageEncoder"),
}
_VAE_MODELS = {
"AutoencoderKLHunyuanVideo": ("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
}
_FAST_VIDEO_MODELS = {
**_TEXT_TO_VIDEO_DIT_MODELS,
**_IMAGE_TO_VIDEO_DIT_MODELS,
**_TEXT_ENCODER_MODELS,
**_IMAGE_ENCODER_MODELS,
**_VAE_MODELS,
}
_SUBPROCESS_COMMAND = [
sys.executable, "-m", "fastvideo.v1.models.dits.registry"
]
_T = TypeVar("_T")
@dataclass(frozen=True)
class _ModelInfo:
architecture: str
@staticmethod
def from_model_cls(model: Type[nn.Module]) -> "_ModelInfo":
return _ModelInfo(
architecture=model.__name__,)
class _BaseRegisteredModel(ABC):
@abstractmethod
def inspect_model_cls(self) -> _ModelInfo:
raise NotImplementedError
@abstractmethod
def load_model_cls(self) -> Type[nn.Module]:
raise NotImplementedError
@dataclass(frozen=True)
class _RegisteredModel(_BaseRegisteredModel):
"""
Represents a model that has already been imported in the main process.
"""
interfaces: _ModelInfo
model_cls: Type[nn.Module]
@staticmethod
def from_model_cls(model_cls: Type[nn.Module]):
return _RegisteredModel(
interfaces=_ModelInfo.from_model_cls(model_cls),
model_cls=model_cls,
)
def inspect_model_cls(self) -> _ModelInfo:
return self.interfaces
def load_model_cls(self) -> Type[nn.Module]:
return self.model_cls
def _run_in_subprocess(fn: Callable[[], _T]) -> _T:
# NOTE: We use a temporary directory instead of a temporary file to avoid
# issues like https://stackoverflow.com/questions/23212435/permission-denied-to-write-to-my-temporary-file
with tempfile.TemporaryDirectory() as tempdir:
output_filepath = os.path.join(tempdir, "registry_output.tmp")
# `cloudpickle` allows pickling lambda functions directly
input_bytes = cloudpickle.dumps((fn, output_filepath))
# cannot use `sys.executable __file__` here because the script
# contains relative imports
returned = subprocess.run(_SUBPROCESS_COMMAND,
input=input_bytes,
capture_output=True)
# check if the subprocess is successful
try:
returned.check_returncode()
except Exception as e:
# wrap raised exception to provide more information
raise RuntimeError(f"Error raised in subprocess:\n"
f"{returned.stderr.decode()}") from e
with open(output_filepath, "rb") as f:
return pickle.load(f)
@dataclass(frozen=True)
class _LazyRegisteredModel(_BaseRegisteredModel):
"""
Represents a model that has not been imported in the main process.
"""
module_name: str
component_name: str
class_name: str
# Performed in another process to avoid initializing CUDA
def inspect_model_cls(self) -> _ModelInfo:
return _run_in_subprocess(
lambda: _ModelInfo.from_model_cls(self.load_model_cls()))
def load_model_cls(self) -> Type[nn.Module]:
mod = importlib.import_module(self.module_name)
return getattr(mod, self.class_name)
@lru_cache(maxsize=128)
def _try_load_model_cls(
model_arch: str,
model: _BaseRegisteredModel,
) -> Optional[Type[nn.Module]]:
from fastvideo.v1.platforms import current_platform
current_platform.verify_model_arch(model_arch)
try:
return model.load_model_cls()
except Exception:
logger.exception("Error in loading model architecture '%s'",
model_arch)
return None
@lru_cache(maxsize=128)
def _try_inspect_model_cls(
model_arch: str,
model: _BaseRegisteredModel,
) -> Optional[_ModelInfo]:
try:
return model.inspect_model_cls()
except Exception:
logger.exception("Error in inspecting model architecture '%s'",
model_arch)
return None
@dataclass
class _ModelRegistry:
# Keyed by model_arch
models: Dict[str, _BaseRegisteredModel] = field(default_factory=dict)
def get_supported_archs(self) -> AbstractSet[str]:
return self.models.keys()
def register_model(
self,
model_arch: str,
model_cls: Union[Type[nn.Module], str],
) -> None:
"""
Register an external model to be used in vLLM.
:code:`model_cls` can be either:
- A :class:`torch.nn.Module` class directly referencing the model.
- A string in the format :code:`<module>:<class>` which can be used to
lazily import the model. This is useful to avoid initializing CUDA
when importing the model and thus the related error
:code:`RuntimeError: Cannot re-initialize CUDA in forked subprocess`.
"""
if model_arch in self.models:
logger.warning(
"Model architecture %s is already registered, and will be "
"overwritten by the new model class %s.", model_arch,
model_cls)
if isinstance(model_cls, str):
split_str = model_cls.split(":")
if len(split_str) != 2:
msg = "Expected a string in the format `<module>:<class>`"
raise ValueError(msg)
model = _LazyRegisteredModel(*split_str)
else:
model = _RegisteredModel.from_model_cls(model_cls)
self.models[model_arch] = model
def _raise_for_unsupported(self, architectures: List[str]):
all_supported_archs = self.get_supported_archs()
if any(arch in all_supported_archs for arch in architectures):
raise ValueError(
f"Model architectures {architectures} failed "
"to be inspected. Please check the logs for more details.")
raise ValueError(
f"Model architectures {architectures} are not supported for now. "
f"Supported architectures: {all_supported_archs}")
def _try_load_model_cls(self,
model_arch: str) -> Optional[Type[nn.Module]]:
if model_arch not in self.models:
return None
return _try_load_model_cls(model_arch, self.models[model_arch])
def _try_inspect_model_cls(self, model_arch: str) -> Optional[_ModelInfo]:
if model_arch not in self.models:
return None
return _try_inspect_model_cls(model_arch, self.models[model_arch])
def _normalize_archs(
self,
architectures: Union[str, List[str]],
) -> List[str]:
if isinstance(architectures, str):
architectures = [architectures]
if not architectures:
logger.warning("No model architectures are specified")
normalized_arch = []
for model in architectures:
if model not in self.models:
model = "TransformersModel"
normalized_arch.append(model)
return normalized_arch
def inspect_model_cls(
self,
architectures: Union[str, List[str]],
) -> Tuple[_ModelInfo, str]:
architectures = self._normalize_archs(architectures)
for arch in architectures:
model_info = self._try_inspect_model_cls(arch)
if model_info is not None:
return (model_info, arch)
return self._raise_for_unsupported(architectures)
def resolve_model_cls(
self,
architectures: Union[str, List[str]],
) -> Tuple[Type[nn.Module], str]:
architectures = self._normalize_archs(architectures)
for arch in architectures:
model_cls = self._try_load_model_cls(arch)
if model_cls is not None:
return (model_cls, arch)
return self._raise_for_unsupported(architectures)
ModelRegistry = _ModelRegistry({
model_arch:
_LazyRegisteredModel(
module_name=f"fastvideo.v1.models.{component_name}.{mod_relname}",
component_name=component_name,
class_name=cls_name,
)
for model_arch, (component_name, mod_relname, cls_name) in _FAST_VIDEO_MODELS.items()
})
@@ -1,239 +0,0 @@
# Copyright 2024 Stability AI, Katherine Crowson 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.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
from dataclasses import dataclass
from typing import Optional, Tuple, Union
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@dataclass
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
Args:
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
denoising loop.
"""
prev_sample: torch.FloatTensor
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
"""
Euler scheduler.
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
methods the library implements for all schedulers such as loading and saving.
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps to train the model.
timestep_spacing (`str`, defaults to `"linspace"`):
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
shift (`float`, defaults to 1.0):
The shift value for the timestep schedule.
reverse (`bool`, defaults to `True`):
Whether to reverse the timestep schedule.
"""
_compatibles = []
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
shift: float = 1.0,
reverse: bool = True,
solver: str = "euler",
n_tokens: Optional[int] = None,
):
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
if not reverse:
sigmas = sigmas.flip(0)
self.sigmas = sigmas
# the value fed to model
self.timesteps = (sigmas[:-1] * num_train_timesteps).to(dtype=torch.float32)
self._step_index = None
self._begin_index = None
self.supported_solver = ["euler"]
if solver not in self.supported_solver:
raise ValueError(f"Solver {solver} not supported. Supported solvers: {self.supported_solver}")
@property
def step_index(self):
"""
The index counter for current timestep. It will increase 1 after each scheduler step.
"""
return self._step_index
@property
def begin_index(self):
"""
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
"""
return self._begin_index
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
def set_begin_index(self, begin_index: int = 0):
"""
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
Args:
begin_index (`int`):
The begin index for the scheduler.
"""
self._begin_index = begin_index
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def set_timesteps(
self,
num_inference_steps: int,
device: Union[str, torch.device] = None,
n_tokens: int = None,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`):
The number of diffusion steps used when generating samples with a pre-trained model.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
"""
self.num_inference_steps = num_inference_steps
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
sigmas = self.sd3_time_shift(sigmas)
if not self.config.reverse:
sigmas = 1 - sigmas
self.sigmas = sigmas
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(dtype=torch.float32, device=device)
# Reset step index
self._step_index = None
def index_for_timestep(self, timestep, schedule_timesteps=None):
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
indices = (schedule_timesteps == timestep).nonzero()
# The sigma index that is taken for the **very** first `step`
# is always the second index (or the last index if there is only 1)
# This way we can ensure we don't accidentally skip a sigma in
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
pos = 1 if len(indices) > 1 else 0
return indices[pos].item()
def _init_step_index(self, timestep):
if self.begin_index is None:
if isinstance(timestep, torch.Tensor):
timestep = timestep.to(self.timesteps.device)
self._step_index = self.index_for_timestep(timestep)
else:
self._step_index = self._begin_index
def scale_model_input(self, sample: torch.Tensor, timestep: Optional[int] = None) -> torch.Tensor:
return sample
def sd3_time_shift(self, t: torch.Tensor):
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
def step(
self,
model_output: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
return_dict: bool = True,
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
process from the learned model outputs (most often the predicted noise).
Args:
model_output (`torch.FloatTensor`):
The direct output from learned diffusion model.
timestep (`float`):
The current discrete timestep in the diffusion chain.
sample (`torch.FloatTensor`):
A current instance of a sample created by the diffusion process.
generator (`torch.Generator`, *optional*):
A random number generator.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
tuple.
Returns:
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
returned, otherwise a tuple is returned where the first element is the sample tensor.
"""
if (isinstance(timestep, int) or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)):
raise ValueError(("Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."), )
if self.step_index is None:
self._init_step_index(timestep)
# Upcast to avoid precision issues when computing prev_sample
sample = sample.to(torch.float32)
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
if self.config.solver == "euler":
prev_sample = sample + model_output.to(torch.float32) * dt
else:
raise ValueError(f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}")
# upon completion increase step index by one
self._step_index += 1
if not return_dict:
return (prev_sample, )
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
def __len__(self):
return self.config.num_train_timesteps
-197
View File
@@ -1,197 +0,0 @@
from dataclasses import dataclass
from typing import Optional, Tuple
import torch
import torch.nn as nn
from transformers.utils import ModelOutput
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def use_default(value, default):
return value if value is not None else default
@dataclass
class TextEncoderModelOutput(ModelOutput):
"""
Base class for model's outputs that also contains a pooling of the last hidden states.
Args:
hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
Sequence of hidden-states at the output of the last layer of the model.
attention_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
Mask to avoid performing attention on padding token indices. Mask values selected in ``[0, 1]``:
hidden_states_list (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed):
Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
text_outputs (`list`, *optional*, returned when `return_texts=True` is passed):
List of decoded texts.
"""
hidden_state: torch.FloatTensor = None
attention_mask: Optional[torch.LongTensor] = None
text_outputs: Optional[list] = None
class TextEncoder(nn.Module):
def __init__(
self,
text_encoder,
tokenizer,
max_length: int,
text_encoder_precision: Optional[str] = None,
text_encoder_path: Optional[str] = None,
output_key: Optional[str] = None,
use_attention_mask: bool = True,
prompt_template: Optional[dict] = None,
prompt_template_video: Optional[dict] = None,
hidden_state_skip_layer: Optional[int] = None,
apply_final_norm: bool = False,
device=None,
):
super().__init__()
# TODO(will): check if there's a cleaner way to do this
self.text_encoder_type = text_encoder.config.architectures[0]
self.max_length = max_length
self.precision = text_encoder_precision
self.model_path = text_encoder_path
self.use_attention_mask = use_attention_mask
if prompt_template_video is not None:
assert (use_attention_mask is True), "Attention mask is True required when training videos."
self.prompt_template = prompt_template
self.prompt_template_video = prompt_template_video
self.hidden_state_skip_layer = hidden_state_skip_layer
self.apply_final_norm = apply_final_norm
if "T5" in self.text_encoder_type:
self.output_key = output_key or "last_hidden_state"
elif "CLIPTextModel" in self.text_encoder_type:
self.output_key = output_key or "pooler_output"
elif "LlamaModel" in self.text_encoder_type or "glm" in self.text_encoder_type:
self.output_key = output_key or "last_hidden_state"
else:
raise ValueError(f"Unsupported text encoder type: {self.text_encoder_type}")
self.model = text_encoder
# self.dtype = self.model.dtype
self.device = device
self.tokenizer = tokenizer
def __repr__(self):
return f"{self.text_encoder_type} ({self.precision} - {self.model_path})"
@staticmethod
def apply_text_to_template(text, template, prevent_empty_text=True):
"""
Apply text to template.
Args:
text (str): Input text.
template (str or list): Template string or list of chat conversation.
prevent_empty_text (bool): If True, we will prevent the user text from being empty
by adding a space. Defaults to True.
"""
if isinstance(template, str):
# Will send string to tokenizer. Used for llm
return template.format(text)
else:
raise TypeError(f"Unsupported template type: {type(template)}")
def text2tokens(self, text):
"""
Tokenize the input text.
Args:
text (str or list): Input text.
"""
if self.prompt_template_video is not None:
prompt_template = self.prompt_template_video["template"]
text = self.apply_text_to_template(text, prompt_template)
kwargs = dict(
truncation=True,
max_length=self.max_length,
padding="max_length",
return_tensors="pt",
)
return self.tokenizer(
text,
return_length=False,
return_overflowing_tokens=False,
return_attention_mask=True,
**kwargs,
)
def encode(
self,
batch_encoding,
use_attention_mask=None,
hidden_state_skip_layer=None,
device=None,
):
"""
Args:
batch_encoding (dict): Batch encoding from tokenizer.
use_attention_mask (bool): Whether to use attention mask. If None, use self.use_attention_mask.
Defaults to None.
output_hidden_states (bool): Whether to output hidden states. If False, return the value of
self.output_key. If True, return the entire output. If set self.hidden_state_skip_layer,
output_hidden_states will be set True. Defaults to False.
hidden_state_skip_layer (int): Number of hidden states to hidden_state_skip_layer. 0 means the last layer.
If None, self.output_key will be used. Defaults to None.
return_texts (bool): Whether to return the decoded texts. Defaults to False.
"""
device = self.model.device if device is None else device
use_attention_mask = use_default(use_attention_mask, self.use_attention_mask)
hidden_state_skip_layer = use_default(hidden_state_skip_layer, self.hidden_state_skip_layer)
attention_mask = (batch_encoding["attention_mask"].to(device) if use_attention_mask else None)
# note: clip will need attention mask
outputs = self.model(
input_ids=batch_encoding["input_ids"].to(device),
attention_mask=attention_mask,
output_hidden_states=hidden_state_skip_layer is not None,
)
if hidden_state_skip_layer is not None:
last_hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
# Real last hidden state already has layer norm applied. So here we only apply it
# for intermediate layers.
if hidden_state_skip_layer > 0 and self.apply_final_norm:
last_hidden_state = self.model.final_layer_norm(last_hidden_state)
else:
last_hidden_state = outputs[self.output_key]
# Remove hidden states of instruction tokens, only keep prompt tokens.
if self.prompt_template_video is not None:
crop_start = self.prompt_template_video.get("crop_start", -1)
last_hidden_state = last_hidden_state[:, crop_start:]
attention_mask = (attention_mask[:, crop_start:] if use_attention_mask else None)
total_length = attention_mask.sum()
last_hidden_state = last_hidden_state[:, :total_length]
return TextEncoderModelOutput(last_hidden_state)
def forward(
self,
text,
use_attention_mask=None,
output_hidden_states=False,
hidden_state_skip_layer=None,
return_texts=False,
):
batch_encoding = self.text2tokens(text)
return self.encode(
batch_encoding,
use_attention_mask=use_attention_mask,
output_hidden_states=output_hidden_states,
hidden_state_skip_layer=hidden_state_skip_layer,
return_texts=return_texts,
)
-97
View File
@@ -1,97 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/utils.py
"""Utils for model executor."""
from typing import Any, Dict, Optional
import torch
# TODO(PY): move it elsewhere
def auto_attributes(init_func):
"""
Decorator that automatically adds all initialization arguments as object attributes.
Example:
@auto_attributes
def __init__(self, a=1, b=2):
pass
# This will automatically set:
# - self.a = 1 and self.b = 2
# - self.config.a = 1 and self.config.b = 2
"""
def wrapper(self, *args, **kwargs):
# Get the function signature
import inspect
signature = inspect.signature(init_func)
parameters = signature.parameters
# Get parameter names (excluding 'self')
param_names = list(parameters.keys())[1:]
# Bind arguments to parameters
bound_args = signature.bind(self, *args, **kwargs)
bound_args.apply_defaults()
# Create config object if it doesn't exist
if not hasattr(self, 'config'):
self.config = type('Config', (), {})()
# Set attributes on self and self.config
for name in param_names:
if name in bound_args.arguments:
value = bound_args.arguments[name]
setattr(self, name, value)
setattr(self.config, name, value)
# Call the original __init__ function
return init_func(self, *args, **kwargs)
return wrapper
def set_random_seed(seed: int) -> None:
from fastvideo.v1.platforms import current_platform
current_platform.seed_everything(seed)
def set_weight_attrs(
weight: torch.Tensor,
weight_attrs: Optional[Dict[str, Any]],
):
"""Set attributes on a weight tensor.
This method is used to set attributes on a weight tensor. This method
will not overwrite existing attributes.
Args:
weight: The weight tensor.
weight_attrs: A dictionary of attributes to set on the weight tensor.
"""
if weight_attrs is None:
return
for key, value in weight_attrs.items():
assert not hasattr(
weight, key), (f"Overwriting existing tensor attribute: {key}")
# NOTE(woosuk): During weight loading, we often do something like:
# narrowed_tensor = param.data.narrow(0, offset, len)
# narrowed_tensor.copy_(real_weight)
# expecting narrowed_tensor and param.data to share the same storage.
# However, on TPUs, narrowed_tensor will lazily propagate to the base
# tensor, which is param.data, leading to the redundant memory usage.
# This sometimes causes OOM errors during model loading. To avoid this,
# we sync the param tensor after its weight loader is called.
# TODO(woosuk): Remove this hack once we have a better solution.
from fastvideo.v1.platforms import current_platform
if current_platform.is_tpu() and key == "weight_loader":
value = _make_synced_weight_loader(value)
setattr(weight, key, value)
def _make_synced_weight_loader(original_weight_loader):
def _synced_weight_loader(param, *args, **kwargs):
original_weight_loader(param, *args, **kwargs)
torch._sync(param)
return _synced_weight_loader
-469
View File
@@ -1,469 +0,0 @@
import torch
from typing import Optional, Tuple
import numpy as np
import torch.nn as nn
from abc import abstractmethod, ABC
import torch.distributed as dist
from fastvideo.v1.distributed import tensor_model_parallel_all_gather, get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size
from math import prod
class ParallelTiledVAE(ABC):
def __init__(self, *args, **kwargs):
# Check if subclass has defined all required properties
required_attributes = [
'tile_sample_min_height',
'tile_sample_min_width',
'tile_sample_min_num_frames',
'tile_sample_stride_height',
'tile_sample_stride_width',
'tile_sample_stride_num_frames',
'use_tiling'
]
for attr in required_attributes:
if not hasattr(self, attr):
raise AttributeError(f"Subclasses of ParallelVAE must define '{attr}' property")
@abstractmethod
def _encode(self, *args, **kwargs):
pass
@abstractmethod
def _decode(self, *args, **kwargs):
pass
def encode(self, x: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, num_frames, height, width = x.shape
if self.use_tiling and num_frames > self.tile_sample_min_num_frames:
latents = self.tiled_encode(x)
elif self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height):
latents = self.spatial_tiled_encode(x)
else:
latents = self._encode(x)
return DiagonalGaussianDistribution(latents)
def decode(self, z: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, num_frames, height, width = z.shape
tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio
tile_latent_min_width = self.tile_sample_stride_width // self.spatial_compression_ratio
tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio
if self.use_tiling and get_sequence_model_parallel_world_size() > 1:
return self.parallel_tiled_decode(z)
if self.use_tiling and num_frames > tile_latent_min_num_frames:
return self.tiled_decode(z)
if self.use_tiling and (width > tile_latent_min_width or height > tile_latent_min_height):
return self.spatial_tiled_decode(z)
return self._decode(z)
def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[-2], b.shape[-2], blend_extent)
for y in range(blend_extent):
b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * (
y / blend_extent
)
return b
def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[-1], b.shape[-1], blend_extent)
for x in range(blend_extent):
b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * (
x / blend_extent
)
return b
def blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[-3], b.shape[-3], blend_extent)
for x in range(blend_extent):
b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (1 - x / blend_extent) + b[:, :, x, :, :] * (
x / blend_extent
)
return b
def spatial_tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
r"""Encode a batch of images using a tiled encoder.
Args:
x (`torch.Tensor`): Input batch of videos.
Returns:
`torch.Tensor`:
The latent representation of the encoded videos.
"""
batch_size, num_channels, num_frames, height, width = x.shape
latent_height = height // self.spatial_compression_ratio
latent_width = width // self.spatial_compression_ratio
tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio
tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio
tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio
tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio
blend_height = tile_latent_min_height - tile_latent_stride_height
blend_width = tile_latent_min_width - tile_latent_stride_width
# Split x into overlapping tiles and encode them separately.
# The tiles have an overlap to avoid seams between tiles.
rows = []
for i in range(0, height, self.tile_sample_stride_height):
row = []
for j in range(0, width, self.tile_sample_stride_width):
tile = x[:, :, :, i : i + self.tile_sample_min_height, j : j + self.tile_sample_min_width]
tile = self._encode(tile)
row.append(tile)
rows.append(row)
return self._merge_spatial_tiles(rows, blend_height, blend_width,tile_latent_stride_height, tile_latent_stride_width)
def _parallel_data_generator(self, gathered_results,
gathered_dim_metadata):
global_idx = 0
for i, per_rank_metadata in enumerate(gathered_dim_metadata):
_start_shape = 0
for shape in per_rank_metadata:
mul_shape = prod(shape)
yield (gathered_results[i, _start_shape:_start_shape +
mul_shape].reshape(shape), global_idx)
_start_shape += mul_shape
global_idx += 1
def parallel_tiled_decode(self, z: torch.FloatTensor) -> torch.FloatTensor:
"""
Parallel version of tiled_decode that distributes both temporal and spatial computation across GPUs
"""
world_size, rank = get_sequence_model_parallel_world_size(), get_sequence_model_parallel_rank()
B, C, T, H, W = z.shape
# Calculate parameters
tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio
tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio
tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio
tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio
tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio
tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio
blend_height = self.tile_sample_min_height - self.tile_sample_stride_height
blend_width = self.tile_sample_min_width - self.tile_sample_stride_width
blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
# Calculate tile dimensions
num_t_tiles = (T + tile_latent_stride_num_frames - 1) // tile_latent_stride_num_frames
num_h_tiles = (H + tile_latent_stride_height - 1) // tile_latent_stride_height
num_w_tiles = (W + tile_latent_stride_width - 1) // tile_latent_stride_width
total_spatial_tiles = num_h_tiles * num_w_tiles
total_tiles = num_t_tiles * total_spatial_tiles
# Calculate tiles per rank and padding
tiles_per_rank = (total_tiles + world_size - 1) // world_size
start_tile_idx = rank * tiles_per_rank
end_tile_idx = min((rank + 1) * tiles_per_rank, total_tiles)
local_results = []
local_dim_metadata = []
# Process assigned tiles
for local_idx, global_idx in enumerate(range(start_tile_idx, end_tile_idx)):
t_idx = global_idx // total_spatial_tiles
spatial_idx = global_idx % total_spatial_tiles
h_idx = spatial_idx // num_w_tiles
w_idx = spatial_idx % num_w_tiles
# Calculate positions
t_start = t_idx * tile_latent_stride_num_frames
h_start = h_idx * tile_latent_stride_height
w_start = w_idx * tile_latent_stride_width
# Extract and process tile
tile = z[:, :, t_start:t_start + tile_latent_min_num_frames + 1,
h_start:h_start + tile_latent_min_height,
w_start:w_start + tile_latent_min_width]
# Process tile
tile = self._decode(tile)
if t_start > 0:
tile = tile[:, :, 1:, :, :]
# Store metadata
shape = tile.shape
# Store decoded data (flattened)
decoded_flat = tile.reshape(-1)
local_results.append(decoded_flat)
local_dim_metadata.append(shape)
results = torch.cat(local_results, dim=0).contiguous()
del local_results
torch.cuda.empty_cache()
# first gather size to pad the results
local_size = torch.tensor([results.size(0)],
device=results.device,
dtype=torch.int64)
all_sizes = [
torch.zeros(1, device=results.device, dtype=torch.int64)
for _ in range(world_size)
]
dist.all_gather(all_sizes, local_size)
max_size = max(size.item() for size in all_sizes)
padded_results = torch.zeros(max_size, device=results.device)
padded_results[:results.size(0)] = results
del results
torch.cuda.empty_cache()
# Gather all results
gathered_dim_metadata = [None] * world_size
gathered_results = torch.zeros_like(padded_results).repeat(
world_size, *[1] * len(padded_results.shape)
).contiguous() # use contiguous to make sure it won't copy data in the following operations
# TODO (PY): use fastvideo distributed methods
dist.all_gather_into_tensor(gathered_results, padded_results)
dist.all_gather_object(gathered_dim_metadata, local_dim_metadata)
# Process gathered results
data = [[[[] for _ in range(num_w_tiles)] for _ in range(num_h_tiles)]
for _ in range(num_t_tiles)]
for current_data, global_idx in self._parallel_data_generator(
gathered_results, gathered_dim_metadata):
t_idx = global_idx // total_spatial_tiles
spatial_idx = global_idx % total_spatial_tiles
h_idx = spatial_idx // num_w_tiles
w_idx = spatial_idx % num_w_tiles
data[t_idx][h_idx][w_idx] = current_data
# Merge results
result_slices = []
last_slice_data = None
for i, tem_data in enumerate(data):
slice_data = self._merge_spatial_tiles(tem_data, blend_height, blend_width,
self.tile_sample_stride_height, self.tile_sample_stride_width)
if i > 0:
slice_data = self.blend_t(last_slice_data, slice_data, blend_num_frames)
result_slices.append(slice_data[:, :, :self.tile_sample_stride_num_frames, :, :])
else:
result_slices.append(slice_data[:, :, :self.tile_sample_stride_num_frames + 1, :, :])
last_slice_data = slice_data
dec = torch.cat(result_slices, dim=2)
return dec
def _merge_spatial_tiles(self, tiles, blend_height, blend_width, stride_height, stride_width):
"""Helper function to merge spatial tiles with blending"""
result_rows = []
for i, row in enumerate(tiles):
result_row = []
for j, tile in enumerate(row):
if i > 0:
tile = self.blend_v(tiles[i - 1][j], tile,
blend_height)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_width)
result_row.append(tile[:, :, :, :stride_height, :stride_width])
result_rows.append(torch.cat(result_row, dim=-1))
return torch.cat(result_rows, dim=-2)
def spatial_tiled_decode(self, z: torch.Tensor) -> torch.Tensor:
r"""
Decode a batch of images using a tiled decoder.
Args:
z (`torch.Tensor`): Input batch of latent vectors.
Returns:
`torch.Tensor`:
The decoded images.
"""
batch_size, num_channels, num_frames, height, width = z.shape
sample_height = height * self.spatial_compression_ratio
sample_width = width * self.spatial_compression_ratio
tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio
tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio
tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio
tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio
blend_height = self.tile_sample_min_height - self.tile_sample_stride_height
blend_width = self.tile_sample_min_width - self.tile_sample_stride_width
# Split z into overlapping tiles and decode them separately.
# The tiles have an overlap to avoid seams between tiles.
rows = []
for i in range(0, height, tile_latent_stride_height):
row = []
for j in range(0, width, tile_latent_stride_width):
tile = z[:, :, :, i : i + tile_latent_min_height, j : j + tile_latent_min_width]
decoded = self._decode(tile)
row.append(decoded)
rows.append(row)
return self._merge_spatial_tiles(rows, blend_height, blend_width,self.tile_sample_stride_height, self.tile_sample_stride_width)
def tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, num_frames, height, width = x.shape
latent_num_frames = (num_frames - 1) // self.temporal_compression_ratio + 1
tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio
tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio
blend_num_frames = tile_latent_min_num_frames - tile_latent_stride_num_frames
row = []
for i in range(0, num_frames, self.tile_sample_stride_num_frames):
tile = x[:, :, i : i + self.tile_sample_min_num_frames + 1, :, :]
if self.use_tiling and (height > self.tile_sample_min_height or width > self.tile_sample_min_width):
tile = self.spatial_tiled_encode(tile)
else:
tile = self._encode(tile)
if i > 0:
tile = tile[:, :, 1:, :, :]
row.append(tile)
result_row = []
for i, tile in enumerate(row):
if i > 0:
tile = self.blend_t(row[i - 1], tile, blend_num_frames)
result_row.append(tile[:, :, :tile_latent_stride_num_frames, :, :])
else:
result_row.append(tile[:, :, : tile_latent_stride_num_frames + 1, :, :])
enc = torch.cat(result_row, dim=2)[:, :, :latent_num_frames]
return enc
def tiled_decode(self, z: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, num_frames, height, width = z.shape
num_sample_frames = (num_frames - 1) * self.temporal_compression_ratio + 1
tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio
tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio
tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio
tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio
blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
row = []
for i in range(0, num_frames, tile_latent_stride_num_frames):
tile = z[:, :, i : i + tile_latent_min_num_frames + 1, :, :]
if self.use_tiling and (tile.shape[-1] > tile_latent_min_width or tile.shape[-2] > tile_latent_min_height):
decoded = self.spatial_tiled_decode(tile)
else:
decoded = self._decode(tile)
if i > 0:
decoded = decoded[:, :, 1:, :, :]
row.append(decoded)
result_row = []
result_row = []
for i, tile in enumerate(row):
if i > 0:
tile = self.blend_t(row[i - 1], tile, blend_num_frames)
result_row.append(tile[:, :, : self.tile_sample_stride_num_frames, :, :])
else:
result_row.append(tile[:, :, : self.tile_sample_stride_num_frames + 1, :, :])
dec = torch.cat(result_row, dim=2)[:, :, :num_sample_frames]
return dec
def enable_tiling(
self,
tile_sample_min_height: Optional[int] = None,
tile_sample_min_width: Optional[int] = None,
tile_sample_min_num_frames: Optional[int] = None,
tile_sample_stride_height: Optional[float] = None,
tile_sample_stride_width: Optional[float] = None,
tile_sample_stride_num_frames: Optional[float] = None,
) -> None:
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
processing larger images.
Args:
tile_sample_min_height (`int`, *optional*):
The minimum height required for a sample to be separated into tiles across the height dimension.
tile_sample_min_width (`int`, *optional*):
The minimum width required for a sample to be separated into tiles across the width dimension.
tile_sample_min_num_frames (`int`, *optional*):
The minimum number of frames required for a sample to be separated into tiles across the frame
dimension.
tile_sample_stride_height (`int`, *optional*):
The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are
no tiling artifacts produced across the height dimension.
tile_sample_stride_width (`int`, *optional*):
The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling
artifacts produced across the width dimension.
tile_sample_stride_num_frames (`int`, *optional*):
The stride between two consecutive frame tiles. This is to ensure that there are no tiling artifacts
produced across the frame dimension.
"""
self.use_tiling = True
self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height
self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width
self.tile_sample_min_num_frames = tile_sample_min_num_frames or self.tile_sample_min_num_frames
self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height
self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width
self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames
def disable_tiling(self) -> None:
r"""
Disable tiled VAE decoding. If `enable_tiling` was previously enabled, this method will go back to computing
decoding in one step.
"""
self.use_tiling = False
# adapted from https://github.com/huggingface/diffusers/blob/e7ffeae0a191f710881d1fbde00cd6ff025e81f2/src/diffusers/models/autoencoders/vae.py#L691
class DiagonalGaussianDistribution(object):
def __init__(self, parameters: torch.Tensor, deterministic: bool = False):
self.parameters = parameters
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
self.deterministic = deterministic
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
if self.deterministic:
self.var = self.std = torch.zeros_like(
self.mean, device=self.parameters.device, dtype=self.parameters.dtype
)
def sample(self, generator: Optional[torch.Generator] = None) -> torch.Tensor:
# make sure sample is on the same device as the parameters and has same dtype
sample = torch.randn(
self.mean.shape,
generator=generator,
device=self.parameters.device,
dtype=self.parameters.dtype,
)
x = self.mean + self.std * sample
return x
def kl(self, other: "DiagonalGaussianDistribution" = None) -> torch.Tensor:
if self.deterministic:
return torch.Tensor([0.0])
else:
if other is None:
return 0.5 * torch.sum(
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar,
dim=[1, 2, 3],
)
else:
return 0.5 * torch.sum(
torch.pow(self.mean - other.mean, 2) / other.var
+ self.var / other.var
- 1.0
- self.logvar
+ other.logvar,
dim=[1, 2, 3],
)
def nll(self, sample: torch.Tensor, dims: Tuple[int, ...] = [1, 2, 3]) -> torch.Tensor:
if self.deterministic:
return torch.Tensor([0.0])
logtwopi = np.log(2.0 * np.pi)
return 0.5 * torch.sum(
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
dim=dims,
)
def mode(self) -> torch.Tensor:
return self.mean
-799
View File
@@ -1,799 +0,0 @@
# Copyright 2024 The Hunyuan Team, The HuggingFace Team and The FastVideo 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.
from typing import Optional, Tuple, Union
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.checkpoint
from fastvideo.v1.models.vaes.common import DiagonalGaussianDistribution, ParallelTiledVAE
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.models.utils import auto_attributes
def prepare_causal_attention_mask(
num_frames: int, height_width: int, dtype: torch.dtype, device: torch.device, batch_size: int = None
) -> torch.Tensor:
indices = torch.arange(1, num_frames + 1, dtype=torch.int32, device=device)
indices_blocks = indices.repeat_interleave(height_width)
x, y = torch.meshgrid(indices_blocks, indices_blocks, indexing="xy")
mask = torch.where(x <= y, 0, -float("inf")).to(dtype=dtype)
if batch_size is not None:
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
return mask
class HunyuanVAEAttention(nn.Module):
def __init__(self, in_channels, heads, dim_head, eps, norm_num_groups, bias):
super().__init__()
self.in_channels = in_channels
self.heads = heads
self.dim_head = dim_head
self.eps = eps
self.norm_num_groups = norm_num_groups
self.bias = bias
inner_dim = heads * dim_head
# Define the projection layers
self.to_q = nn.Linear(in_channels, inner_dim, bias=bias)
self.to_k = nn.Linear(in_channels, inner_dim, bias=bias)
self.to_v = nn.Linear(in_channels, inner_dim, bias=bias)
self.to_out = nn.Sequential(
nn.Linear(inner_dim, in_channels, bias=bias)
)
# Optional normalization layers
self.group_norm = nn.GroupNorm(norm_num_groups, in_channels, eps=eps, affine=True)
def forward(self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
residual = hidden_states
batch_size, sequence_length, _ = hidden_states.shape
hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
# Project to query, key, value
query = self.to_q(hidden_states)
key = self.to_k(hidden_states)
value = self.to_v(hidden_states)
# Reshape for multi-head attention
head_dim = self.dim_head
query = query.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
# Perform scaled dot-product attention
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
# Reshape back
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.heads * head_dim)
hidden_states = hidden_states.to(query.dtype)
# Linear projection
hidden_states = self.to_out(hidden_states)
# Residual connection and rescale
hidden_states = hidden_states + residual
return hidden_states
class HunyuanVideoCausalConv3d(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: Union[int, Tuple[int, int, int]] = 3,
stride: Union[int, Tuple[int, int, int]] = 1,
padding: Union[int, Tuple[int, int, int]] = 0,
dilation: Union[int, Tuple[int, int, int]] = 1,
bias: bool = True,
pad_mode: str = "replicate",
) -> None:
super().__init__()
kernel_size = (kernel_size, kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size
self.pad_mode = pad_mode
self.time_causal_padding = (
kernel_size[0] // 2,
kernel_size[0] // 2,
kernel_size[1] // 2,
kernel_size[1] // 2,
kernel_size[2] - 1,
0,
)
self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, dilation, bias=bias)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = F.pad(hidden_states, self.time_causal_padding, mode=self.pad_mode)
return self.conv(hidden_states)
class HunyuanVideoUpsampleCausal3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: Optional[int] = None,
kernel_size: int = 3,
stride: int = 1,
bias: bool = True,
upsample_factor: Tuple[float, float, float] = (2, 2, 2),
) -> None:
super().__init__()
out_channels = out_channels or in_channels
self.upsample_factor = upsample_factor
self.conv = HunyuanVideoCausalConv3d(in_channels, out_channels, kernel_size, stride, bias=bias)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
num_frames = hidden_states.size(2)
first_frame, other_frames = hidden_states.split((1, num_frames - 1), dim=2)
first_frame = F.interpolate(
first_frame.squeeze(2), scale_factor=self.upsample_factor[1:], mode="nearest"
).unsqueeze(2)
if num_frames > 1:
# See: https://github.com/pytorch/pytorch/issues/81665
# Unless you have a version of pytorch where non-contiguous implementation of F.interpolate
# is fixed, this will raise either a runtime error, or fail silently with bad outputs.
# If you are encountering an error here, make sure to try running encoding/decoding with
# `vae.enable_tiling()` first. If that doesn't work, open an issue at:
# https://github.com/huggingface/diffusers/issues
other_frames = other_frames.contiguous()
other_frames = F.interpolate(other_frames, scale_factor=self.upsample_factor, mode="nearest")
hidden_states = torch.cat((first_frame, other_frames), dim=2)
else:
hidden_states = first_frame
hidden_states = self.conv(hidden_states)
return hidden_states
class HunyuanVideoDownsampleCausal3D(nn.Module):
def __init__(
self,
channels: int,
out_channels: Optional[int] = None,
padding: int = 1,
kernel_size: int = 3,
bias: bool = True,
stride=2,
) -> None:
super().__init__()
out_channels = out_channels or channels
self.conv = HunyuanVideoCausalConv3d(channels, out_channels, kernel_size, stride, padding, bias=bias)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.conv(hidden_states)
return hidden_states
class HunyuanVideoResnetBlockCausal3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: Optional[int] = None,
dropout: float = 0.0,
groups: int = 32,
eps: float = 1e-6,
non_linearity: str = "silu",
) -> None:
super().__init__()
out_channels = out_channels or in_channels
self.nonlinearity = get_act_fn(non_linearity)
self.norm1 = nn.GroupNorm(groups, in_channels, eps=eps, affine=True)
self.conv1 = HunyuanVideoCausalConv3d(in_channels, out_channels, 3, 1, 0)
self.norm2 = nn.GroupNorm(groups, out_channels, eps=eps, affine=True)
self.dropout = nn.Dropout(dropout)
self.conv2 = HunyuanVideoCausalConv3d(out_channels, out_channels, 3, 1, 0)
self.conv_shortcut = None
if in_channels != out_channels:
self.conv_shortcut = HunyuanVideoCausalConv3d(in_channels, out_channels, 1, 1, 0)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = hidden_states.contiguous()
residual = hidden_states
hidden_states = self.norm1(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.conv1(hidden_states)
hidden_states = self.norm2(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.dropout(hidden_states)
hidden_states = self.conv2(hidden_states)
if self.conv_shortcut is not None:
residual = self.conv_shortcut(residual)
hidden_states = hidden_states + residual
return hidden_states
class HunyuanVideoMidBlock3D(nn.Module):
def __init__(
self,
in_channels: int,
dropout: float = 0.0,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_act_fn: str = "silu",
resnet_groups: int = 32,
add_attention: bool = True,
attention_head_dim: int = 1,
) -> None:
super().__init__()
resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)
self.add_attention = add_attention
# There is always at least one resnet
resnets = [
HunyuanVideoResnetBlockCausal3D(
in_channels=in_channels,
out_channels=in_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
non_linearity=resnet_act_fn,
)
]
attentions = []
for _ in range(num_layers):
if self.add_attention:
attentions.append(
HunyuanVAEAttention(
in_channels,
heads=in_channels // attention_head_dim,
dim_head=attention_head_dim,
eps=resnet_eps,
norm_num_groups=resnet_groups,
bias=True,
)
)
else:
attentions.append(None)
resnets.append(
HunyuanVideoResnetBlockCausal3D(
in_channels=in_channels,
out_channels=in_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
non_linearity=resnet_act_fn,
)
)
self.attentions = nn.ModuleList(attentions)
self.resnets = nn.ModuleList(resnets)
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(self.resnets[0], hidden_states)
for attn, resnet in zip(self.attentions, self.resnets[1:]):
if attn is not None:
batch_size, num_channels, num_frames, height, width = hidden_states.shape
hidden_states = hidden_states.permute(0, 2, 3, 4, 1).flatten(1, 3)
attention_mask = prepare_causal_attention_mask(
num_frames, height * width, hidden_states.dtype, hidden_states.device, batch_size=batch_size
)
hidden_states = attn(hidden_states, attention_mask=attention_mask)
hidden_states = hidden_states.unflatten(1, (num_frames, height, width)).permute(0, 4, 1, 2, 3)
hidden_states = self._gradient_checkpointing_func(resnet, hidden_states)
else:
hidden_states = self.resnets[0](hidden_states)
for attn, resnet in zip(self.attentions, self.resnets[1:]):
if attn is not None:
batch_size, num_channels, num_frames, height, width = hidden_states.shape
hidden_states = hidden_states.permute(0, 2, 3, 4, 1).flatten(1, 3)
attention_mask = prepare_causal_attention_mask(
num_frames, height * width, hidden_states.dtype, hidden_states.device, batch_size=batch_size
)
hidden_states = attn(hidden_states, attention_mask=attention_mask)
hidden_states = hidden_states.unflatten(1, (num_frames, height, width)).permute(0, 4, 1, 2, 3)
hidden_states = resnet(hidden_states)
return hidden_states
class HunyuanVideoDownBlock3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
dropout: float = 0.0,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_act_fn: str = "silu",
resnet_groups: int = 32,
add_downsample: bool = True,
downsample_stride: int = 2,
downsample_padding: int = 1,
) -> None:
super().__init__()
resnets = []
for i in range(num_layers):
in_channels = in_channels if i == 0 else out_channels
resnets.append(
HunyuanVideoResnetBlockCausal3D(
in_channels=in_channels,
out_channels=out_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
non_linearity=resnet_act_fn,
)
)
self.resnets = nn.ModuleList(resnets)
if add_downsample:
self.downsamplers = nn.ModuleList(
[
HunyuanVideoDownsampleCausal3D(
out_channels,
out_channels=out_channels,
padding=downsample_padding,
stride=downsample_stride,
)
]
)
else:
self.downsamplers = None
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
if torch.is_grad_enabled() and self.gradient_checkpointing:
for resnet in self.resnets:
hidden_states = self._gradient_checkpointing_func(resnet, hidden_states)
else:
for resnet in self.resnets:
hidden_states = resnet(hidden_states)
if self.downsamplers is not None:
for downsampler in self.downsamplers:
hidden_states = downsampler(hidden_states)
return hidden_states
class HunyuanVideoUpBlock3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
dropout: float = 0.0,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_act_fn: str = "silu",
resnet_groups: int = 32,
add_upsample: bool = True,
upsample_scale_factor: Tuple[int, int, int] = (2, 2, 2),
) -> None:
super().__init__()
resnets = []
for i in range(num_layers):
input_channels = in_channels if i == 0 else out_channels
resnets.append(
HunyuanVideoResnetBlockCausal3D(
in_channels=input_channels,
out_channels=out_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
non_linearity=resnet_act_fn,
)
)
self.resnets = nn.ModuleList(resnets)
if add_upsample:
self.upsamplers = nn.ModuleList(
[
HunyuanVideoUpsampleCausal3D(
out_channels,
out_channels=out_channels,
upsample_factor=upsample_scale_factor,
)
]
)
else:
self.upsamplers = None
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
if torch.is_grad_enabled() and self.gradient_checkpointing:
for resnet in self.resnets:
hidden_states = self._gradient_checkpointing_func(resnet, hidden_states)
else:
for resnet in self.resnets:
hidden_states = resnet(hidden_states)
if self.upsamplers is not None:
for upsampler in self.upsamplers:
hidden_states = upsampler(hidden_states)
return hidden_states
class HunyuanVideoEncoder3D(nn.Module):
r"""
Causal encoder for 3D video-like data introduced in [Hunyuan Video](https://huggingface.co/papers/2412.03603).
"""
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: Tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
),
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
double_z: bool = True,
mid_block_add_attention=True,
temporal_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
) -> None:
super().__init__()
self.conv_in = HunyuanVideoCausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
if down_block_type != "HunyuanVideoDownBlock3D":
raise ValueError(f"Unsupported down_block_type: {down_block_type}")
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(temporal_compression_ratio))
if temporal_compression_ratio == 4:
add_spatial_downsample = bool(i < num_spatial_downsample_layers)
add_time_downsample = bool(
i >= (len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block
)
elif temporal_compression_ratio == 8:
add_spatial_downsample = bool(i < num_spatial_downsample_layers)
add_time_downsample = bool(i < num_time_downsample_layers)
else:
raise ValueError(f"Unsupported time_compression_ratio: {temporal_compression_ratio}")
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2,) if add_time_downsample else (1,)
downsample_stride = tuple(downsample_stride_T + downsample_stride_HW)
down_block = HunyuanVideoDownBlock3D(
num_layers=layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(add_spatial_downsample or add_time_downsample),
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
downsample_stride=downsample_stride,
downsample_padding=0,
)
self.down_blocks.append(down_block)
self.mid_block = HunyuanVideoMidBlock3D(
in_channels=block_out_channels[-1],
resnet_eps=1e-6,
resnet_act_fn=act_fn,
attention_head_dim=block_out_channels[-1],
resnet_groups=norm_num_groups,
add_attention=mid_block_add_attention,
)
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6)
self.conv_act = nn.SiLU()
conv_out_channels = 2 * out_channels if double_z else out_channels
self.conv_out = HunyuanVideoCausalConv3d(block_out_channels[-1], conv_out_channels, kernel_size=3)
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.conv_in(hidden_states)
if torch.is_grad_enabled() and self.gradient_checkpointing:
for down_block in self.down_blocks:
hidden_states = self._gradient_checkpointing_func(down_block, hidden_states)
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
else:
for down_block in self.down_blocks:
hidden_states = down_block(hidden_states)
hidden_states = self.mid_block(hidden_states)
hidden_states = self.conv_norm_out(hidden_states)
hidden_states = self.conv_act(hidden_states)
hidden_states = self.conv_out(hidden_states)
return hidden_states
class HunyuanVideoDecoder3D(nn.Module):
r"""
Causal decoder for 3D video-like data introduced in [Hunyuan Video](https://huggingface.co/papers/2412.03603).
"""
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
up_block_types: Tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
),
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
mid_block_add_attention=True,
time_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = HunyuanVideoCausalConv3d(in_channels, block_out_channels[-1], kernel_size=3, stride=1)
self.up_blocks = nn.ModuleList([])
# mid
self.mid_block = HunyuanVideoMidBlock3D(
in_channels=block_out_channels[-1],
resnet_eps=1e-6,
resnet_act_fn=act_fn,
attention_head_dim=block_out_channels[-1],
resnet_groups=norm_num_groups,
add_attention=mid_block_add_attention,
)
# up
reversed_block_out_channels = list(reversed(block_out_channels))
output_channel = reversed_block_out_channels[0]
for i, up_block_type in enumerate(up_block_types):
if up_block_type != "HunyuanVideoUpBlock3D":
raise ValueError(f"Unsupported up_block_type: {up_block_type}")
prev_output_channel = output_channel
output_channel = reversed_block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_upsample_layers = int(np.log2(spatial_compression_ratio))
num_time_upsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_upsample = bool(i < num_spatial_upsample_layers)
add_time_upsample = bool(
i >= len(block_out_channels) - 1 - num_time_upsample_layers and not is_final_block
)
else:
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}")
upsample_scale_factor_HW = (2, 2) if add_spatial_upsample else (1, 1)
upsample_scale_factor_T = (2,) if add_time_upsample else (1,)
upsample_scale_factor = tuple(upsample_scale_factor_T + upsample_scale_factor_HW)
up_block = HunyuanVideoUpBlock3D(
num_layers=self.layers_per_block + 1,
in_channels=prev_output_channel,
out_channels=output_channel,
add_upsample=bool(add_spatial_upsample or add_time_upsample),
upsample_scale_factor=upsample_scale_factor,
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
)
self.up_blocks.append(up_block)
prev_output_channel = output_channel
# out
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6)
self.conv_act = nn.SiLU()
self.conv_out = HunyuanVideoCausalConv3d(block_out_channels[0], out_channels, kernel_size=3)
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.conv_in(hidden_states)
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
for up_block in self.up_blocks:
hidden_states = self._gradient_checkpointing_func(up_block, hidden_states)
else:
hidden_states = self.mid_block(hidden_states)
for up_block in self.up_blocks:
hidden_states = up_block(hidden_states)
# post-process
hidden_states = self.conv_norm_out(hidden_states)
hidden_states = self.conv_act(hidden_states)
hidden_states = self.conv_out(hidden_states)
return hidden_states
class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
r"""
A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos.
Introduced in [HunyuanVideo](https://huggingface.co/papers/2412.03603).
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
for all models (such as downloading or saving).
"""
_supports_gradient_checkpointing = True
@auto_attributes
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
latent_channels: int = 16,
down_block_types: Tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
),
up_block_types: Tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
),
block_out_channels: Tuple[int] = (128, 256, 512, 512),
layers_per_block: int = 2,
act_fn: str = "silu",
norm_num_groups: int = 32,
scaling_factor: float = 0.476986,
spatial_compression_ratio: int = 8,
temporal_compression_ratio: int = 4,
mid_block_add_attention: bool = True,
load_encoder: bool = True,
load_decoder: bool = True,
) -> None:
super().__init__()
# TODO(will): only pass in config. We do this by manually defining a
# config for hunyuan vae
self.block_out_channels = block_out_channels
if load_encoder:
self.encoder = HunyuanVideoEncoder3D(
in_channels=in_channels,
out_channels=latent_channels,
down_block_types=down_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
double_z=True,
mid_block_add_attention=mid_block_add_attention,
temporal_compression_ratio=temporal_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
)
self.quant_conv = nn.Conv3d(2 * latent_channels, 2 * latent_channels, kernel_size=1)
if load_decoder:
self.decoder = HunyuanVideoDecoder3D(
in_channels=latent_channels,
out_channels=out_channels,
up_block_types=up_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
time_compression_ratio=temporal_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=mid_block_add_attention,
)
self.post_quant_conv = nn.Conv3d(latent_channels, latent_channels, kernel_size=1)
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
# intermediate tiles together, the memory requirement can be lowered.
self.use_tiling = True
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 16
# The minimal distance between two spatial tiles
self.tile_sample_stride_height = 192
self.tile_sample_stride_width = 192
self.tile_sample_stride_num_frames = 12
def _encode(self, x: torch.Tensor) -> torch.Tensor:
x = self.encoder(x)
enc = self.quant_conv(x)
return enc
def _decode(self, z: torch.Tensor) -> torch.Tensor:
z = self.post_quant_conv(z)
dec = self.decoder(z)
return dec
def forward(
self,
sample: torch.Tensor,
sample_posterior: bool = False,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
r"""
Args:
sample (`torch.Tensor`): Input sample.
sample_posterior (`bool`, *optional*, defaults to `False`):
Whether to sample from the posterior.
"""
x = sample
posterior = self.encode(x).latent_dist
if sample_posterior:
z = posterior.sample(generator=generator)
else:
z = posterior.mode()
dec = self.decode(z)
return dec
-260
View File
@@ -1,260 +0,0 @@
"""
Diffusion pipelines for fastvideo.v1.
This package contains diffusion pipelines for generating videos and images.
"""
import os
import json
from copy import deepcopy
from typing import Dict, Optional, Type, Any
from fastvideo.v1.pipelines.pipeline_registry import PipelineRegistry
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.logger import init_logger
from huggingface_hub import snapshot_download
from transformers import PretrainedConfig
from fastvideo.v1.models.hf_transformer_utils import get_hf_config, get_diffusers_config
from fastvideo.v1.models import get_scheduler
import glob
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
logger = init_logger(__name__)
import tempfile
import filelock
import hashlib
# Then import the base classes
from fastvideo.v1.pipelines.composed.composed_pipeline_base import (
ComposedPipelineBase,
DiffusionPipelineOutput
)
def get_pipeline_type(inference_args: InferenceArgs) -> str:
# hardcode for now
return "hunyuan_video"
def get_lock(model_name_or_path: str):
lock_dir = tempfile.gettempdir()
os.makedirs(os.path.dirname(lock_dir), exist_ok=True)
model_name = model_name_or_path.replace("/", "-")
hash_name = hashlib.sha256(model_name.encode()).hexdigest()
# add hash to avoid conflict with old users' lock files
lock_file_name = hash_name + model_name + ".lock"
# mode 0o666 is required for the filelock to be shared across users
lock = filelock.FileLock(os.path.join(lock_dir, lock_file_name),
mode=0o666)
return lock
def maybe_download_model(model_path: str) -> str:
"""
Check if the model path is a Hugging Face Hub model ID and download it if needed.
Args:
model_path: Local path or Hugging Face Hub model ID
Returns:
Local path to the model
"""
# If the path exists locally, return it
if os.path.exists(model_path):
logger.info(f"Model already exists locally at {model_path}")
return model_path
# Otherwise, assume it's a HF Hub model ID and try to download it
try:
logger.info(f"Downloading model snapshot from HF Hub for {model_path}...")
with get_lock(model_path):
local_path = snapshot_download(
repo_id=model_path,
ignore_patterns=["*.onnx", "*.msgpack"],
)
logger.info(f"Downloaded model to {local_path}")
return local_path
except Exception as e:
raise ValueError(f"Could not find model at {model_path} and failed to download from HF Hub: {e}")
def verify_model_config_and_directory(model_path: str) -> dict:
"""
Verify that the model directory contains a valid diffusers configuration.
Args:
model_path: Path to the model directory
Returns:
The loaded model configuration as a dictionary
"""
# Check for model_index.json which is required for diffusers models
config_path = os.path.join(model_path, "model_index.json")
if not os.path.exists(config_path):
raise ValueError(
f"Model directory {model_path} does not contain model_index.json. "
"Only Hugging Face diffusers format is supported."
)
# Check for transformer and vae directories
transformer_dir = os.path.join(model_path, "transformer")
vae_dir = os.path.join(model_path, "vae")
if not os.path.exists(transformer_dir):
raise ValueError(f"Model directory {model_path} does not contain a transformer/ directory.")
if not os.path.exists(vae_dir):
raise ValueError(f"Model directory {model_path} does not contain a vae/ directory.")
# Load the config
try:
with open(config_path, "r") as f:
config = json.load(f)
except Exception as e:
raise ValueError(f"Failed to load model configuration from {config_path}: {e}")
# Verify diffusers version exists
if "_diffusers_version" not in config:
raise ValueError(f"model_index.json does not contain _diffusers_version")
logger.info(f"Diffusers version: {config['_diffusers_version']}")
return config
def load_pipeline_module(module_name: str, component_model_path: str, transformers_or_diffusers: str, architecture: str, inference_args: InferenceArgs) -> Any:
"""
Load a pipeline module using the appropriate loader.
Args:
module_name: Name of the module (e.g., "vae", "text_encoder", "transformer", "scheduler")
component_model_path: Path to the component model
transformers_or_diffusers: Whether the module is from transformers or diffusers
architecture: Architecture of the component model
inference_args: Inference arguments
Returns:
The loaded module
"""
return PipelineComponentLoader.load_module(
module_name=module_name,
component_model_path=component_model_path,
transformers_or_diffusers=transformers_or_diffusers,
architecture=architecture,
inference_args=inference_args
)
def load_pipeline_modules(model_path: str, config: Dict, inference_args: InferenceArgs) -> dict[str, Any]:
"""
Load the pipeline modules from the config.
Args:
config: The model_index.json config
inference_args: Inference arguments
Returns:
Dictionary mapping module names to loaded modules
"""
logger.info(f"Loading pipeline modules from config: {config}")
modules_config = deepcopy(config)
# remove keys that are not pipeline modules
modules_config.pop("_class_name")
modules_config.pop("_diffusers_version")
# some sanity checks
assert len(modules_config) > 1, "model_index.json must contain at least one pipeline module"
required_modules = ["vae", "text_encoder", "transformer", "scheduler", "tokenizer"]
for module_name in required_modules:
if module_name not in modules_config:
raise ValueError(f"model_index.json must contain a {module_name} module")
logger.info(f"Diffusers config passed sanity checks")
# all the component models used by the pipeline
pipeline_modules = {}
for module_name, (transformers_or_diffusers, architecture) in modules_config.items():
component_model_path = os.path.join(model_path, module_name)
module = load_pipeline_module(
module_name,
component_model_path,
transformers_or_diffusers,
architecture,
inference_args,
)
pipeline_modules[module_name] = module
# Check if all required modules were loaded
for module_name in required_modules:
if module_name not in pipeline_modules or pipeline_modules[module_name] is None:
logger.warning(f"Required module {module_name} was not loaded properly")
return pipeline_modules
def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
"""
Only works with valid hf diffusers configs. (model_index.json)
We want to build a pipeline based on the inference args mode_path:
1. download the model from the hub if it's not already downloaded
2. verify the model config and directory
3. based on the config, determine the pipeline class
4. parse the config to get the model components (vae, text_encoders, etc...)
5. the pipeline loader class will use the model component names and paths to load
6. the pipeline class will be composed of the models returned by the pipeline loader
"""
# Get pipeline type
model_path = inference_args.model_path
model_path = maybe_download_model(model_path)
# inference_args.downloaded_model_path = model_path
logger.info(f"Model path: {model_path}")
config = verify_model_config_and_directory(model_path)
pipeline_architecture = config.get("_class_name")
if pipeline_architecture is None:
raise ValueError("Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
pipeline_cls, pipeline_architecture = PipelineRegistry.resolve_pipeline_cls(pipeline_architecture)
# instantiate the pipeline
pipeline = pipeline_cls()
pipeline_modules = load_pipeline_modules(model_path, config, inference_args)
logger.info(f"Initializing encoders")
pipeline.initialize_encoders(pipeline_modules, inference_args)
logger.info(f"Registering modules")
pipeline.register_modules(pipeline_modules)
logger.info(f"Setting up pipeline")
pipeline.setup_pipeline(inference_args)
logger.info(f"Initializing pipeline")
pipeline.initialize_pipeline(inference_args)
# pipeline is now initialized and ready to use
return pipeline
def list_available_pipelines() -> Dict[str, Type[Any]]:
"""
List all available pipeline types.
Returns:
A dictionary of pipeline names to pipeline classes.
"""
return PipelineRegistry.list()
__all__ = [
"build_pipeline",
"list_available_pipelines",
"ComposedPipelineBase",
"DiffusionPipelineOutput",
]
@@ -1,18 +0,0 @@
"""
Composed pipelines for diffusion models.
This package contains pipelines that are composed of multiple stages.
"""
from fastvideo.v1.pipelines.composed.composed_pipeline_base import (
ComposedPipelineBase,
DiffusionPipelineOutput,
)
__all__ = [
"ComposedPipelineBase",
"DiffusionPipelineOutput",
]
# Note: Do not import TextToVideoPipeline here to avoid circular imports.
# It will be imported in the main pipelines/__init__.py file.
@@ -1,147 +0,0 @@
"""
Base class for composed pipelines.
This module defines the base class for pipelines that are composed of multiple stages.
"""
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Union, Tuple
import torch
import numpy as np
from tqdm.auto import tqdm
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.pipelines.stages import PipelineStage
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@dataclass
class DiffusionPipelineOutput:
"""Output from a diffusion pipeline."""
videos: Union[torch.Tensor, np.ndarray]
class ComposedPipelineBase(ABC):
"""
Base class for pipelines composed of multiple stages.
This class provides the framework for creating pipelines by composing multiple
stages together. Each stage is responsible for a specific part of the diffusion
process, and the pipeline orchestrates the execution of these stages.
"""
is_video_pipeline: bool = False # To be overridden by video pipelines
def __init__(self):
"""
Initialize the pipeline.
The pipeline should be completely stateless and not hold any batch
state.
"""
self._stages: List[PipelineStage] = []
self._modules: Dict[str, Any] = {}
self._stage_name_mapping: Dict[str, PipelineStage] = {}
@property
def device(self) -> torch.device:
"""Get the device for this pipeline."""
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
@property
def modules(self) -> Dict[str, Any]:
"""Get all modules used by this pipeline."""
return self._modules
@abstractmethod
def setup_pipeline(self, inference_args: InferenceArgs):
"""
Setup the pipeline.
"""
...
@abstractmethod
def initialize_encoders(self, modules: Dict[str, Any], inference_args: InferenceArgs):
"""
Initialize the encoders. Will remove the encoders/tokenizers modules from the
modules. Will add the TextEncoder or ImageEncoder to the modules.
"""
...
def register_modules(self, modules: Dict[str, Any]):
"""
Register modules with the pipeline and its stages.
We will use the _module_name_mapping to map the module names used
in the Diffusers config to the internal names (how it can be accessed
in the pipeline).
Args:
modules: The modules to register.
"""
self._modules.update(modules)
# Register modules with self
for name, module in modules.items():
setattr(self, name, module)
# self._modules[name] = module
# Register modules with stages that need them
for stage in self._stages:
stage.register_modules(modules)
# TODO(will): perhaps we should not register all modules with the
# stage. See below.
# stage_modules = {}
# for name, module in mapped_modules.items():
# if hasattr(stage, f"needs_{name}") and getattr(stage, f"needs_{name}"):
# stage_modules[name] = module
# if stage_modules:
# stage.register_modules(**stage_modules)
def add_stage(self, name: str, stage: PipelineStage):
assert self._modules is not None, "No modules are registered"
stage.register_modules(self._modules)
self._stages.append(stage)
self._stage_name_mapping[name] = stage
setattr(self, name, stage)
# TODO(will): don't hardcode no_grad
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
) -> DiffusionPipelineOutput:
"""
Generate a video or image using the pipeline.
Args:
prompt: The prompt(s) to guide generation.
negative_prompt: The negative prompt(s) to guide generation.
height: The height of the generated video/image.
width: The width of the generated video/image.
num_frames: The number of frames to generate (for video).
num_inference_steps: The number of inference steps.
guidance_scale: The scale for classifier-free guidance.
num_videos_per_prompt: The number of videos to generate per prompt.
generator: The random number generator.
latents: The initial latents.
output_type: The output type.
**kwargs: Additional arguments.
Returns:
The generated video or image.
"""
# Execute each stage
for stage in self._stages:
batch = stage(batch, inference_args)
# Return the output
return DiffusionPipelineOutput(videos=batch.output)
@@ -1,13 +0,0 @@
"""
HunYuan pipeline implementations.
This package contains implementations of diffusion pipelines for HunYuan models.
"""
from fastvideo.v1.pipelines.implementations.hunyuan.hunyuan_pipeline import (
EntryClass,
)
__all__ = [
"EntryClass",
]
@@ -1,90 +0,0 @@
import os
import torch
__all__ = [
"PROMPT_TEMPLATE",
"MODEL_BASE",
"PRECISIONS",
"NORMALIZATION_TYPE",
"ACTIVATION_TYPE",
"VAE_PATH",
"TEXT_ENCODER_PATH",
"TOKENIZER_PATH",
"TEXT_PROJECTION",
"DATA_TYPE",
"NEGATIVE_PROMPT",
]
PRECISION_TO_TYPE = {
"fp32": torch.float32,
"fp16": torch.float16,
"bf16": torch.bfloat16,
}
# =================== Constant Values =====================
# Computation scale factor, 1P = 1_000_000_000_000_000. Tensorboard will display the value in PetaFLOPS to avoid
# overflow error when tensorboard logging values.
# When using decoder-only models, we must provide a prompt template to instruct the text encoder
# on how to generate the text.
# --------------------------------------------------------------------
# TODO: Many models will have this model specific prompt template. Create a centralized place to store them.
# TODO: add inference arg: --use-template to use the model specific prompt template. Else use no template
PROMPT_TEMPLATE_ENCODE = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
PROMPT_TEMPLATE_ENCODE_VIDEO = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
PROMPT_TEMPLATE = {
"image": {
"template": PROMPT_TEMPLATE_ENCODE,
"crop_start": 36,
},
"video": {
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
"crop_start": 95,
},
}
# ======================= Model ======================
PRECISIONS = {"fp32", "fp16", "bf16"}
NORMALIZATION_TYPE = {"layer", "rms"}
ACTIVATION_TYPE = {"relu", "silu", "gelu", "gelu_tanh"}
# =================== Model Path =====================
MODEL_BASE = os.getenv("MODEL_BASE", "./data/hunyuan")
# =================== Data =======================
DATA_TYPE = {"image", "video", "image_video"}
# 3D VAE
VAE_PATH = {"884-16c-hy": f"{MODEL_BASE}/hunyuan-video-t2v-720p/vae"}
# Text Encoder
TEXT_ENCODER_PATH = {
"clipL": f"{MODEL_BASE}/text_encoder_2",
"llm": f"{MODEL_BASE}/text_encoder",
}
# Tokenizer
TOKENIZER_PATH = {
"clipL": f"{MODEL_BASE}/text_encoder_2",
"llm": f"{MODEL_BASE}/text_encoder",
}
TEXT_PROJECTION = {
"linear", # Default, an nn.Linear() layer
"single_refiner", # Single TokenRefiner. Refer to LI-DiT
}
@@ -1,186 +0,0 @@
"""
HunYuan video diffusion pipeline implementation.
This module contains an implementation of the HunYuan video diffusion pipeline
using the modular pipeline architecture.
"""
from typing import Union, Any, Dict
import torch
from diffusers.image_processor import VaeImageProcessor
from fastvideo.v1.pipelines.composed import ComposedPipelineBase
from fastvideo.v1.pipelines.stages import (
InputValidationStage,
PromptEncodingStage,
TimestepPreparationStage,
LatentPreparationStage,
ConditioningStage,
DenoisingStage,
DecodingStage,
PostProcessingStage,
)
from fastvideo.v1.pipelines.stages.prompt_encoding import PromptEncodingStage
# from fastvideo.v1.pipelines.stages.timestep_preparation import FlowMatchingTimestepPreparationStage
from fastvideo.v1.inference_args import InferenceArgs
# from fastvideo.v1.pipelines.composed.composed_pipeline_base import DiffusionPipelineOutput
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
# TODO(will): move PRECISION_TO_TYPE to better place
from .constants import PROMPT_TEMPLATE, PRECISION_TO_TYPE
from diffusers.utils import BaseOutput
import numpy as np
from dataclasses import dataclass
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
# class HunyuanLatentPreparationStage(LatentPreparationStage):
# def _call_implementation(self, batch: ForwardBatch, inference_args: InferenceArgs) -> ForwardBatch:
# "custom logic for HunYuan latent preparation"
# pass
# class hunyuanloader(PipelineLoader):
# def load_components(self, inference_args: InferenceArgs):
# pass
@dataclass
class DiffusionPipelineOutput(BaseOutput):
videos: Union[torch.Tensor, np.ndarray]
class HunyuanVideoPipeline(ComposedPipelineBase):
def initialize_encoders(self, modules: Dict[str, Any], inference_args: InferenceArgs):
self.initialize_encoders_v1(modules, inference_args)
def initialize_encoders_v1(self, modules: Dict[str, Any], inference_args: InferenceArgs):
"""
Initialize the encoders. Will remove the encoders/tokenizers modules from the
modules. Will add the TextEncoder or ImageEncoder to the modules.
"""
from fastvideo.v1.models.text_encoder import TextEncoder
crop_start = PROMPT_TEMPLATE["video"].get("crop_start", 0)
max_length = inference_args.text_len + crop_start
# prompt_template
prompt_template = PROMPT_TEMPLATE["image"]
# prompt_template_video
prompt_template_video = PROMPT_TEMPLATE["video"]
encoder_1 = modules.pop("text_encoder")
assert encoder_1 is not None, "Text encoder is not found"
encoder_1.to(inference_args.device)
encoder_1.to(dtype=PRECISION_TO_TYPE[inference_args.text_encoder_precision])
encoder_1.requires_grad_(False)
tokenizer_1 = modules.pop("tokenizer")
assert tokenizer_1 is not None, "Tokenizer is not found"
text_encoder = TextEncoder(
text_encoder=encoder_1,
tokenizer=tokenizer_1,
# text_encoder_type="text_encoder"
max_length=max_length,
# text_encoder_precision=inference_args.text_encoder_precision,
prompt_template=prompt_template,
prompt_template_video=prompt_template_video,
hidden_state_skip_layer=inference_args.hidden_state_skip_layer,
apply_final_norm=False,
device=inference_args.device if not inference_args.use_cpu_offload else "cpu",
)
encoder_2 = modules.pop("text_encoder_2")
assert encoder_2 is not None, "Text encoder 2 is not found"
encoder_2.to(inference_args.device)
encoder_2.to(dtype=PRECISION_TO_TYPE[inference_args.text_encoder_precision])
encoder_2.requires_grad_(False)
tokenizer_2 = modules.pop("tokenizer_2")
assert tokenizer_2 is not None, "Tokenizer 2 is not found"
text_encoder_2 = TextEncoder(
text_encoder=encoder_2,
tokenizer=tokenizer_2,
# text_encoder_type="text_encoder_2",
max_length=inference_args.text_len_2,
# text_encoder_precision=inference_args.text_encoder_precision,
device=inference_args.device if not inference_args.use_cpu_offload else "cpu",
)
modules["text_encoder"] = text_encoder
modules["text_encoder_2"] = text_encoder_2
def setup_pipeline(self, inference_args: InferenceArgs):
self.add_stage("input_validation_stage",
InputValidationStage())
self.add_stage("prompt_encoding_stage_primary",
PromptEncodingStage(is_secondary=False))
self.add_stage("prompt_encoding_stage_secondary",
PromptEncodingStage(is_secondary=True))
self.add_stage("conditioning_stage",
ConditioningStage())
self.add_stage("timestep_preparation_stage",
TimestepPreparationStage())
self.add_stage("latent_preparation_stage",
LatentPreparationStage())
self.add_stage("denoising_stage",
DenoisingStage())
self.add_stage("decoding_stage",
DecodingStage())
def initialize_pipeline(self, inference_args: InferenceArgs):
assert len(self._stages) > 0, "Pipeline stages are not set"
assert len(self._modules) > 0, "Pipeline modules are not set"
vae_scale_factor = 2**(len(self.vae.block_out_channels) - 1)
inference_args.vae_scale_factor = vae_scale_factor
self.image_processor = VaeImageProcessor(vae_scale_factor=vae_scale_factor)
self.register_modules({"image_processor": self.image_processor})
num_channels_latents = self.transformer.in_channels
inference_args.num_channels_latents = num_channels_latents
# TODO
def adjust_video_length(self, batch: ForwardBatch, inference_args: InferenceArgs):
"""Adjust video length based on VAE version"""
video_length = batch.num_frames
batch.num_frames = (video_length - 1) // 4 + 1
return batch
@torch.no_grad()
def forward(self, batch: ForwardBatch, inference_args: InferenceArgs) -> ForwardBatch:
logger.info(f"Running pipeline stages: {self._stage_name_mapping.keys()}")
logger.info(f"Batch: {batch}")
# for stage in self._stages:
# batch = stage(batch, inference_args)
# or
batch = self.input_validation_stage(batch, inference_args)
batch = self.prompt_encoding_stage_primary(batch, inference_args)
batch = self.prompt_encoding_stage_secondary(batch, inference_args)
batch = self.conditioning_stage(batch, inference_args)
batch = self.timestep_preparation_stage(batch, inference_args)
# custom logic
batch = self.adjust_video_length(batch, inference_args)
batch = self.latent_preparation_stage(batch, inference_args)
batch = self.denoising_stage(batch, inference_args)
batch = self.decoding_stage(batch, inference_args)
return DiffusionPipelineOutput(videos=batch.videos)
EntryClass = HunyuanVideoPipeline
@@ -1,105 +0,0 @@
"""
Data structures for functional pipeline processing.
This module defines the dataclasses used to pass state between pipeline components
in a functional manner, reducing the need for explicit parameter passing.
"""
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Union, Tuple, Callable
import torch
@dataclass
class ForwardBatch:
"""
Complete state passed through the pipeline execution.
This dataclass contains all information needed during the diffusion pipeline
execution, allowing methods to update specific components without needing
to manage numerous individual parameters.
"""
# TODO(will): double check that args are separate from inference_args
# properly. Also maybe think about providing an abstraction for pipeline
# specific arguments.
data_type: str
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None
mask_strategy: Optional[Dict[str, List[str]]] = None
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
negative_prompt: Optional[Union[str, List[str]]] = None
# Primary encoder embeddings
prompt_embeds: Optional[torch.Tensor] = None
negative_prompt_embeds: Optional[torch.Tensor] = None
attention_mask: Optional[torch.Tensor] = None
negative_attention_mask: Optional[torch.Tensor] = None
# Secondary encoder embeddings (for dual-encoder models)
prompt_embeds_2: Optional[torch.Tensor] = None
negative_prompt_embeds_2: Optional[torch.Tensor] = None
attention_mask_2: Optional[torch.Tensor] = None
negative_attention_mask_2: Optional[torch.Tensor] = None
# Additional text-related parameters
max_sequence_length: Optional[int] = None
prompt_template: Optional[Dict[str, Any]] = None
do_classifier_free_guidance: bool = False
# Batch info
batch_size: Optional[int] = None
num_videos_per_prompt: int = 1
# Tracking if embeddings are already processed
is_prompt_processed: bool = False
# Latent tensors
latents: Optional[torch.Tensor] = None
noise_pred: Optional[torch.Tensor] = None
# Latent dimensions
num_channels_latents: Optional[int] = None
height_latents: Optional[int] = None
width_latents: Optional[int] = None
num_frames: int = 1 # Default for image models
# Original dimensions (before VAE scaling)
height: Optional[int] = None
width: Optional[int] = None
# Timesteps
timesteps: Optional[torch.Tensor] = None
timestep: Optional[Union[torch.Tensor, float, int]] = None
step_index: Optional[int] = None
# Scheduler parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
eta: float = 0.0
sigmas: Optional[List[float]] = None
n_tokens: Optional[int] = None
# Other parameters that may be needed by specific schedulers
extra_step_kwargs: Dict[str, Any] = field(default_factory=dict)
# Component modules (populated by the pipeline)
modules: Dict[str, Any] = field(default_factory=dict)
# Final output (after pipeline completion)
output: Any = None
# Extra parameters that might be needed by specific pipeline implementations
extra: Dict[str, Any] = field(default_factory=dict)
device: torch.device = field(default_factory=lambda: torch.device("cuda"))
def __post_init__(self):
"""Initialize dependent fields after dataclass initialization."""
# Set do_classifier_free_guidance based on guidance scale and negative prompt
if self.guidance_scale > 1.0:
self.do_classifier_free_guidance = True
@@ -1,91 +0,0 @@
# Adapted from https://github.com/vllm-project/vllm/blob/v0.6.4.post1/vllm/model_executor/models/registry.py
# and https://github.com/sgl-project/sglang/blob/v0.4.3/python/sglang/srt/models/registry.py
import importlib
import pkgutil
from dataclasses import dataclass, field
from functools import lru_cache
from typing import AbstractSet, Dict, List, Optional, Tuple, Type, Union
from fastvideo.v1.logger import init_logger
import torch.nn as nn
logger = init_logger(__name__)
@dataclass
class _PipelineRegistry:
# Keyed by pipeline_arch
pipelines: Dict[str, Union[Type[nn.Module], str]] = field(default_factory=dict)
def get_supported_archs(self) -> AbstractSet[str]:
return self.pipelines.keys()
def _raise_for_unsupported(self, architectures: List[str]):
all_supported_archs = self.get_supported_archs()
if any(arch in all_supported_archs for arch in architectures):
raise ValueError(
f"Pipeline architectures {architectures} failed "
"to be inspected. Please check the logs for more details."
)
raise ValueError(
f"Pipeline architectures {architectures} are not supported for now. "
f"Supported architectures: {all_supported_archs}"
)
def _try_load_pipeline_cls(self, pipeline_arch: str) -> Optional[Type[nn.Module]]:
if pipeline_arch not in self.pipelines:
return None
return self.pipelines[pipeline_arch]
def resolve_pipeline_cls(
self,
architecture: str,
) -> Tuple[Type[nn.Module], str]:
if not architecture:
logger.warning("No pipeline architecture is specified")
pipeline_cls = self._try_load_pipeline_cls(architecture)
if pipeline_cls is not None:
return (pipeline_cls, architecture)
return self._raise_for_unsupported(architecture)
@lru_cache()
def import_pipeline_classes():
pipeline_arch_name_to_cls = {}
package_name = "fastvideo.v1.pipelines.implementations"
package = importlib.import_module(package_name)
for _, name, ispkg in pkgutil.iter_modules(package.__path__, package_name + "."):
if ispkg:
try:
module = importlib.import_module(name)
except Exception as e:
logger.warning(f"Ignore import error when loading {name}. " f"{e}")
continue
if hasattr(module, "EntryClass"):
entry = module.EntryClass
print(entry)
print(entry.__name__)
if isinstance(
entry, list
): # To support multiple pipeline classes in one module
for tmp in entry:
assert (
tmp.__name__ not in pipeline_arch_name_to_cls
), f"Duplicated pipeline implementation for {tmp.__name__}"
pipeline_arch_name_to_cls[tmp.__name__] = tmp
else:
assert (
entry.__name__ not in pipeline_arch_name_to_cls
), f"Duplicated pipeline implementation for {entry.__name__}"
pipeline_arch_name_to_cls[entry.__name__] = entry
return pipeline_arch_name_to_cls
PipelineRegistry = _PipelineRegistry(import_pipeline_classes())
-28
View File
@@ -1,28 +0,0 @@
"""
Pipeline stages for diffusion models.
This package contains the various stages that can be composed to create
complete diffusion pipelines.
"""
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.input_validation import InputValidationStage
from fastvideo.v1.pipelines.stages.prompt_encoding import PromptEncodingStage
from fastvideo.v1.pipelines.stages.timestep_preparation import TimestepPreparationStage
from fastvideo.v1.pipelines.stages.latent_preparation import LatentPreparationStage
from fastvideo.v1.pipelines.stages.conditioning import ConditioningStage
from fastvideo.v1.pipelines.stages.denoising import DenoisingStage
from fastvideo.v1.pipelines.stages.decoding import DecodingStage
from fastvideo.v1.pipelines.stages.post_processing import PostProcessingStage
__all__ = [
"PipelineStage",
"InputValidationStage",
"PromptEncodingStage",
"TimestepPreparationStage",
"LatentPreparationStage",
"ConditioningStage",
"DenoisingStage",
"DecodingStage",
"PostProcessingStage",
]
-125
View File
@@ -1,125 +0,0 @@
"""
Base classes for pipeline stages.
This module defines the abstract base classes for pipeline stages that can be
composed to create complete diffusion pipelines.
"""
from abc import ABC, abstractmethod
import torch
import time
import traceback
from typing import Dict, Any
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class PipelineStage(ABC):
"""
Abstract base class for all pipeline stages.
A pipeline stage represents a discrete step in the diffusion process that can be
composed with other stages to create a complete pipeline. Each stage is responsible
for a specific part of the process, such as prompt encoding, latent preparation, etc.
"""
def __init__(self, enable_logging: bool = False):
"""
Initialize the pipeline stage.
Args:
enable_logging: Whether to enable logging for this stage.
"""
self._enable_logging = enable_logging
self._stage_name = self.__class__.__name__
self._logger = init_logger(f"fastvideo.v1.pipelines.stages.{self._stage_name}")
@property
def device(self) -> torch.device:
"""Get the device for this stage."""
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
def set_logging(self, enable: bool):
"""
Enable or disable logging for this stage.
Args:
enable: Whether to enable logging.
"""
self._enable_logging = enable
def __call__(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
) -> ForwardBatch:
"""
Execute the stage's processing on the batch with optional logging.
Should not be overridden by subclasses.
Args:
batch: The current batch information.
inference_args: The inference arguments.
Returns:
The updated batch information after this stage's processing.
"""
if self._enable_logging:
self._logger.info(f"[{self._stage_name}] Starting execution")
start_time = time.time()
try:
# Call the actual implementation
result = self._call_implementation(batch, inference_args)
execution_time = time.time() - start_time
self._logger.info(f"[{self._stage_name}] Execution completed in {execution_time * 1000:.2f} ms")
return result
except Exception as e:
execution_time = time.time() - start_time
self._logger.error(f"[{self._stage_name}] Error during execution after {execution_time * 1000:.2f} ms: {e}")
self._logger.error(f"[{self._stage_name}] Traceback: {traceback.format_exc()}")
# Re-raise the exception
raise
else:
# Just call the implementation directly if logging is disabled
return self._call_implementation(batch, inference_args)
@abstractmethod
def _call_implementation(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
) -> ForwardBatch:
"""
Actual implementation of the stage's processing.
This method should be implemented by subclasses to provide the actual
processing logic for the stage.
Args:
batch: The current batch information.
inference_args: The inference arguments.
Returns:
The updated batch information after this stage's processing.
"""
pass
def register_modules(self, modules: Dict[str, Any]):
"""
Register modules needed by this stage.
Args:
modules: The modules to register.
"""
for name, module in modules.items():
if self._enable_logging:
self._logger.debug(f"[{self._stage_name}] Registering module: {name}")
setattr(self, name, module)
@@ -1,70 +0,0 @@
"""
Conditioning stage for diffusion pipelines.
"""
import torch
from typing import Optional
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class ConditioningStage(PipelineStage):
"""
Stage for applying conditioning to the diffusion process.
This stage handles the application of conditioning, such as classifier-free guidance,
to the diffusion process.
"""
def _call_implementation(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
) -> ForwardBatch:
"""
Apply conditioning to the diffusion process.
Args:
batch: The current batch information.
inference_args: The inference arguments.
Returns:
The batch with applied conditioning.
"""
if not batch.do_classifier_free_guidance:
return batch
logger.info(f"batch.negative_prompt_embeds: {batch.negative_prompt_embeds}")
logger.info(f"do_classifier_free_guidance: {batch.do_classifier_free_guidance}")
logger.info(f"cfg_scale: {batch.guidance_scale}")
# Ensure negative prompt embeddings are available
assert batch.negative_prompt_embeds is not None, (
"Negative prompt embeddings are required for classifier-free guidance"
)
# Concatenate primary embeddings and masks
batch.prompt_embeds = torch.cat(
[batch.negative_prompt_embeds, batch.prompt_embeds]
)
if batch.attention_mask is not None:
batch.attention_mask = torch.cat(
[batch.negative_attention_mask, batch.attention_mask]
)
# Concatenate secondary embeddings and masks if present
if batch.prompt_embeds_2 is not None:
batch.prompt_embeds_2 = torch.cat(
[batch.negative_prompt_embeds_2, batch.prompt_embeds_2]
)
if batch.attention_mask_2 is not None:
batch.attention_mask_2 = torch.cat(
[batch.negative_attention_mask_2, batch.attention_mask_2]
)
return batch
-77
View File
@@ -1,77 +0,0 @@
"""
Decoding stage for diffusion pipelines.
"""
import torch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.utils import PRECISION_TO_TYPE
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class DecodingStage(PipelineStage):
"""
Stage for decoding latent representations into pixel space.
This stage handles the decoding of latent representations into the final
output format (e.g., pixel values).
"""
def _call_implementation(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
) -> ForwardBatch:
"""
Decode latent representations into pixel space.
Args:
batch: The current batch information.
inference_args: The inference arguments.
Returns:
The batch with decoded outputs.
"""
latents = batch.latents
# Skip decoding if output type is latent
if inference_args.output_type == "latent":
image = latents
else:
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
vae_autocast_enabled = (vae_dtype != torch.float32) and not inference_args.disable_autocast
# Apply scaling/shifting if needed
if (hasattr(self.vae.config, "shift_factor") and self.vae.config.shift_factor):
latents = (latents / self.vae.config.scaling_factor + self.vae.config.shift_factor)
else:
latents = latents / self.vae.config.scaling_factor
# Decode latents
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
if inference_args.vae_tiling:
self.vae.enable_tiling()
# if inference_args.vae_sp:
# self.vae.enable_parallel()
image = self.vae.decode(latents)
# Normalize image to [0, 1] range
image = (image / 2 + 0.5).clamp(0, 1)
# Convert to CPU float32 for compatibility
image = image.cpu().float()
# Update batch with decoded image
batch.videos = image
# Offload models if needed
if hasattr(self, 'maybe_free_model_hooks'):
self.maybe_free_model_hooks()
return batch
-210
View File
@@ -1,210 +0,0 @@
"""
Denoising stage for diffusion pipelines.
"""
import inspect
import torch
from typing import Optional, Dict, Any, List
from tqdm.auto import tqdm
from einops import rearrange
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.utils import PRECISION_TO_TYPE
# TODO(will-refactor): change this to fastvideo.distributed
from fastvideo.v1.distributed import get_sequence_model_parallel_world_size, get_sequence_model_parallel_rank
from fastvideo.v1.distributed.communication_op import sequence_model_parallel_all_gather
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class DenoisingStage(PipelineStage):
"""
Stage for running the denoising loop in diffusion pipelines.
This stage handles the iterative denoising process that transforms
the initial noise into the final output.
"""
def _call_implementation(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
) -> ForwardBatch:
"""
Run the denoising loop.
Args:
batch: The current batch information.
inference_args: The inference arguments.
Returns:
The batch with denoised latents.
"""
# Prepare extra step kwargs for scheduler
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta
},
)
# Setup precision and autocast settings
target_dtype = PRECISION_TO_TYPE[inference_args.precision]
autocast_enabled = (target_dtype != torch.float32) and not inference_args.disable_autocast
# Handle sequence parallelism if enabled
world_size, rank = get_sequence_model_parallel_world_size(), get_sequence_model_parallel_rank()
sp_group = True if world_size > 1 else False
if sp_group:
latents = rearrange(batch.latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
latents = latents[:, :, rank, :, :, :]
batch.latents = latents
# Get timesteps and calculate warmup steps
timesteps = batch.timesteps
num_inference_steps = batch.num_inference_steps
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
# Create 3D list for mask strategy
def dict_to_3d_list(mask_strategy, t_max=50, l_max=60, h_max=24):
result = [[[None for _ in range(h_max)] for _ in range(l_max)] for _ in range(t_max)]
if mask_strategy is None:
return result
for key, value in mask_strategy.items():
t, l, h = map(int, key.split('_'))
result[t][l][h] = value
return result
mask_strategy = dict_to_3d_list(batch.mask_strategy)
# Get latents and embeddings
latents = batch.latents
prompt_embeds = batch.prompt_embeds
prompt_embeds_2 = batch.prompt_embeds_2
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
# Skip if interrupted
if hasattr(self, 'interrupt') and self.interrupt:
continue
# Expand latents for classifier-free guidance
latent_model_input = (torch.cat([latents] * 2) if batch.do_classifier_free_guidance else latents)
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
# Prepare inputs for transformer
t_expand = t.repeat(latent_model_input.shape[0])
guidance_expand = (torch.tensor(
[inference_args.embedded_cfg_scale] * latent_model_input.shape[0],
dtype=torch.float32,
device=batch.device,
).to(target_dtype) * 1000.0 if inference_args.embedded_cfg_scale is not None else None)
# Predict noise residual
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
# Prepare encoder hidden states
if prompt_embeds_2 is not None and prompt_embeds_2.shape[-1] != prompt_embeds.shape[-1]:
prompt_embeds_2 = torch.nn.functional.pad(
prompt_embeds_2,
(0, prompt_embeds.shape[2] - prompt_embeds_2.shape[1]),
value=0,
).unsqueeze(1)
encoder_hidden_states = torch.cat([prompt_embeds_2, prompt_embeds], dim=1) if prompt_embeds_2 is not None else prompt_embeds
# Run transformer
noise_pred = self.transformer(
latent_model_input,
encoder_hidden_states,
t_expand,
guidance=guidance_expand,
)
# Apply guidance
if batch.do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + batch.guidance_scale * (noise_pred_text - noise_pred_uncond)
# Apply guidance rescale if needed
if batch.guidance_rescale > 0.0:
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
noise_pred = self.rescale_noise_cfg(
noise_pred,
noise_pred_text,
guidance_rescale=batch.guidance_rescale,
)
# Compute the previous noisy sample
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
# Update progress bar
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
if progress_bar is not None:
progress_bar.update()
# Gather results if using sequence parallelism
if sp_group:
latents = sequence_model_parallel_all_gather(latents, dim=2)
# Update batch with final latents
batch.latents = latents
return batch
def prepare_extra_func_kwargs(self, func, kwargs):
"""
Prepare extra kwargs for the scheduler step.
Args:
func: The function to prepare kwargs for.
kwargs: The kwargs to prepare.
Returns:
The prepared kwargs.
"""
extra_step_kwargs = {}
for k, v in kwargs.items():
accepts = k in set(inspect.signature(func).parameters.keys())
if accepts:
extra_step_kwargs[k] = v
return extra_step_kwargs
def progress_bar(self, iterable=None, total=None):
"""
Create a progress bar for the denoising process.
Args:
iterable: The iterable to iterate over.
total: The total number of items.
Returns:
A tqdm progress bar.
"""
return tqdm(iterable=iterable, total=total)
def rescale_noise_cfg(self, noise_cfg, noise_pred_text, guidance_rescale=0.0):
"""
Rescale noise prediction according to guidance_rescale.
Based on findings of "Common Diffusion Noise Schedules and Sample Steps are Flawed"
(https://arxiv.org/pdf/2305.08891.pdf), Section 3.4.
Args:
noise_cfg: The noise prediction with guidance.
noise_pred_text: The text-conditioned noise prediction.
guidance_rescale: The guidance rescale factor.
Returns:
The rescaled noise prediction.
"""
std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True)
std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)
# Rescale the results from guidance (fixes overexposure)
noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
# Mix with the original results from guidance by factor guidance_rescale
noise_cfg = (guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg)
return noise_cfg
@@ -1,91 +0,0 @@
"""
Input validation stage for diffusion pipelines.
"""
from typing import Optional, Union, List
import torch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.logger import init_logger
import random
logger = init_logger(__name__)
class InputValidationStage(PipelineStage):
"""
Stage for validating and preparing inputs for diffusion pipelines.
This stage validates that all required inputs are present and properly formatted
before proceeding with the diffusion process.
"""
def _generate_seeds(self, batch: ForwardBatch, inference_args: InferenceArgs):
"""Generate seeds for the inference"""
seed = inference_args.seed
num_videos_per_prompt = inference_args.num_videos
seeds = [seed + i for i in range(num_videos_per_prompt)]
batch.seeds = seeds
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
batch.generator = [torch.Generator("cpu").manual_seed(seed) for seed in seeds]
def _call_implementation(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
) -> ForwardBatch:
"""
Validate and prepare inputs.
Args:
batch: The current batch information.
inference_args: The inference arguments.
Returns:
The validated batch information.
"""
self._generate_seeds(batch, inference_args)
# Ensure prompt is properly formatted
if batch.prompt is None and batch.prompt_embeds is None:
raise ValueError("Either `prompt` or `prompt_embeds` must be provided")
# Ensure negative prompt is properly formatted if using classifier-free guidance
if batch.do_classifier_free_guidance:
if batch.negative_prompt is None and batch.negative_prompt_embeds is None:
raise ValueError(
"For classifier-free guidance, either `negative_prompt` or "
"`negative_prompt_embeds` must be provided"
)
# Validate height and width
if batch.height % 8 != 0 or batch.width % 8 != 0:
raise ValueError(
f"Height and width must be divisible by 8 but are {batch.height} and {batch.width}."
)
# Validate number of inference steps
if batch.num_inference_steps <= 0:
raise ValueError(
f"Number of inference steps must be positive, but got {batch.num_inference_steps}"
)
# Validate guidance scale if using classifier-free guidance
if batch.do_classifier_free_guidance and batch.guidance_scale <= 0:
raise ValueError(
f"Guidance scale must be positive, but got {batch.guidance_scale}"
)
# Set device if not already set
if batch.device is None:
batch.device = self.device
# Set data type if not already set
if batch.data_type is None:
batch.data_type = inference_args.precision
return batch
@@ -1,113 +0,0 @@
"""
Latent preparation stage for diffusion pipelines.
"""
import torch
from typing import Optional, Tuple
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.inference_args import InferenceArgs
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class LatentPreparationStage(PipelineStage):
"""
Stage for preparing initial latent variables for the diffusion process.
This stage handles the preparation of the initial latent variables that will be
denoised during the diffusion process.
"""
def __init__(self):
super().__init__()
# TODO(will): this is a hack to get the vae scale factor. Check if this
# is only needed for hunyuan
def _call_implementation(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
) -> ForwardBatch:
"""
Prepare initial latent variables for the diffusion process.
Args:
batch: The current batch information.
inference_args: The inference arguments.
Returns:
The batch with prepared latent variables.
"""
# Determine batch size
if isinstance(batch.prompt, list):
batch_size = len(batch.prompt)
elif batch.prompt is not None:
batch_size = 1
else:
batch_size = batch.prompt_embeds.shape[0]
# Adjust batch size for number of videos per prompt
batch_size *= batch.num_videos_per_prompt
# Get required parameters
dtype = batch.prompt_embeds.dtype
device = batch.device
generator = batch.generator
latents = batch.latents
num_frames = batch.num_frames
height = batch.height
width = batch.width
# Calculate latent shape
shape = (
batch_size,
inference_args.num_channels_latents,
num_frames,
int(height) // inference_args.vae_scale_factor,
int(width) // inference_args.vae_scale_factor,
)
# Validate generator if it's a list
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."
)
# Generate or use provided latents
if latents is None:
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
else:
latents = latents.to(device)
# Scale the initial noise if needed
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
# Update batch with prepared latents
batch.latents = latents
# Adjust video length based on VAE version if needed
if hasattr(self, 'adjust_video_length'):
batch = self.adjust_video_length(batch, inference_args)
return batch
def adjust_video_length(self, batch: ForwardBatch, inference_args: InferenceArgs) -> ForwardBatch:
"""
Adjust video length based on VAE version.
Args:
batch: The current batch information.
inference_args: The inference arguments.
Returns:
The batch with adjusted video length.
"""
video_length = batch.num_frames
# TODO
batch.num_frames = (video_length - 1) // 4 + 1
return batch
@@ -1,61 +0,0 @@
"""
Post-processing stage for diffusion pipelines.
"""
import torch
import numpy as np
from typing import Optional, Union, List, Tuple
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.inference_args import InferenceArgs
from diffusers.utils import BaseOutput
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class PostProcessingStage(PipelineStage):
"""
Stage for post-processing the decoded outputs.
This stage handles any final processing needed on the decoded outputs,
such as format conversion, normalization, etc.
"""
def _call_implementation(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
) -> ForwardBatch:
"""
Apply post-processing to the results.
Args:
batch: The current batch information.
inference_args: The inference arguments.
Returns:
The batch with post-processed outputs.
"""
videos = batch.videos
# Convert to numpy if requested
if inference_args.output_type == "numpy":
videos = videos.numpy()
# Create output object
output = DiffusionPipelineOutput(videos=videos)
batch.output = output
return batch
class DiffusionPipelineOutput(BaseOutput):
"""
Output class for diffusion pipelines.
Args:
videos: The generated videos.
"""
videos: Union[torch.Tensor, np.ndarray]
@@ -1,113 +0,0 @@
"""
Prompt encoding stages for diffusion pipelines.
This module contains implementations of prompt encoding stages for diffusion pipelines.
"""
from typing import Any, Dict, List, Optional, Union, Tuple
import torch
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.pipelines.stages import PipelineStage
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class PromptEncodingStage(PipelineStage):
"""
Stage for encoding text prompts into embeddings for diffusion models.
This stage handles the encoding of text prompts into the embedding space
expected by the diffusion model.
"""
def __init__(self, enable_logging: bool = False, is_secondary: bool = False):
"""
Initialize the prompt encoding stage.
Args:
enable_logging: Whether to enable logging for this stage.
is_secondary: Whether this is a secondary text encoder.
"""
super().__init__(enable_logging=enable_logging)
self.is_secondary = is_secondary
def _call_implementation(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
) -> ForwardBatch:
"""
Encode the prompt into text encoder hidden states.
Args:
batch: The current batch information.
inference_args: The inference arguments.
Returns:
The batch with encoded prompt embeddings.
"""
if self.is_secondary:
assert self.text_encoder_2 is not None, "Secondary text encoder is not set"
text_encoder = self.text_encoder_2
else:
text_encoder = self.text_encoder
prompt: Union[str, List[str]] = batch.prompt
device: torch.device = batch.device
num_videos_per_prompt: int = batch.num_videos_per_prompt
data_type: str = batch.data_type
# Get the right prompt embeds and attention masks based on whether this is primary or secondary
if self.is_secondary:
prompt_embeds = batch.prompt_embeds_2
attention_mask = batch.attention_mask_2
else:
prompt_embeds = batch.prompt_embeds
attention_mask = batch.attention_mask
if prompt_embeds is None:
# textual inversion: process multi-vector tokens if necessary
# if isinstance(self, TextualInversionLoaderMixin):
# prompt = self.maybe_convert_prompt(prompt, text_encoder.tokenizer)
text_inputs = text_encoder.text2tokens(prompt)
prompt_outputs = text_encoder.encode(text_inputs, device=device)
prompt_embeds = prompt_outputs.hidden_state
# TODO(will): support clip_skip
if text_encoder is not None:
# TODO(will-refactor): use text_encoder.dtype
prompt_embeds_dtype = torch.float16
elif self.transformer is not None:
prompt_embeds_dtype = self.transformer.dtype
else:
prompt_embeds_dtype = prompt_embeds.dtype
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device)
# print("prompt_embeds", type(prompt_embeds))
# logger.info(f"prompt_embeds shape: {prompt_embeds.shape}")
if prompt_embeds.ndim == 1:
prompt_embeds = prompt_embeds.unsqueeze(0)
if prompt_embeds.ndim == 2:
bs_embed, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt)
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, -1)
else:
bs_embed, seq_len, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, seq_len, -1)
# Set the appropriate attributes based on whether this is primary or secondary
if self.is_secondary:
batch.prompt_embeds_2 = prompt_embeds
else:
batch.prompt_embeds = prompt_embeds
return batch
@@ -1,86 +0,0 @@
"""
Timestep preparation stages for diffusion pipelines.
This module contains implementations of timestep preparation stages for diffusion pipelines.
"""
from typing import Any, Dict, List, Optional, Union, Tuple
import torch
import inspect
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class TimestepPreparationStage(PipelineStage):
"""
Stage for preparing timesteps for the diffusion process.
This stage handles the preparation of the timestep sequence that will be used
during the diffusion process.
"""
def _call_implementation(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
) -> ForwardBatch:
"""
Prepare timesteps for the diffusion process.
Args:
batch: The current batch information.
inference_args: The inference arguments.
Returns:
The batch with prepared timesteps.
"""
scheduler = self.scheduler
device = batch.device
num_inference_steps = batch.num_inference_steps
timesteps = batch.timesteps
sigmas = batch.sigmas
n_tokens = batch.n_tokens
# Prepare extra kwargs for set_timesteps
extra_set_timesteps_kwargs = {}
if n_tokens is not None and "n_tokens" in inspect.signature(scheduler.set_timesteps).parameters:
extra_set_timesteps_kwargs["n_tokens"] = n_tokens
# Handle custom timesteps or sigmas
if timesteps is not None and sigmas is not None:
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
if timesteps is not None:
accepts_timesteps = "timesteps" in inspect.signature(scheduler.set_timesteps).parameters
if not accepts_timesteps:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" timestep schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(timesteps=timesteps, device=device, **extra_set_timesteps_kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
elif sigmas is not None:
accept_sigmas = "sigmas" in inspect.signature(scheduler.set_timesteps).parameters
if not accept_sigmas:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" sigmas schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(sigmas=sigmas, device=device, **extra_set_timesteps_kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
else:
scheduler.set_timesteps(num_inference_steps, device=device, **extra_set_timesteps_kwargs)
timesteps = scheduler.timesteps
# Update batch with prepared timesteps
batch.timesteps = timesteps
batch.num_inference_steps = num_inference_steps
return batch
-102
View File
@@ -1,102 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import logging
import traceback
from contextlib import suppress
from itertools import chain
from typing import TYPE_CHECKING, Optional
from vllm.plugins import load_plugins_by_group
from vllm.utils import resolve_obj_by_qualname
from .interface import _Backend # noqa: F401
from .interface import Platform, PlatformEnum
logger = logging.getLogger(__name__)
def cuda_platform_plugin() -> Optional[str]:
is_cuda = False
try:
from vllm.utils import import_pynvml
pynvml = import_pynvml()
pynvml.nvmlInit()
try:
# NOTE: Edge case: vllm cpu build on a GPU machine.
# Third-party pynvml can be imported in cpu build,
# we need to check if vllm is built with cpu too.
# Otherwise, vllm will always activate cuda plugin
# on a GPU machine, even if in a cpu build.
is_cuda = (pynvml.nvmlDeviceGetCount() > 0)
finally:
pynvml.nvmlShutdown()
except Exception as e:
if "nvml" not in e.__class__.__name__.lower():
# If the error is not related to NVML, re-raise it.
raise e
# CUDA is supported on Jetson, but NVML may not be.
import os
def cuda_is_jetson() -> bool:
return os.path.isfile("/etc/nv_tegra_release") \
or os.path.exists("/sys/class/tegra-firmware")
if cuda_is_jetson():
is_cuda = True
return "vllm.platforms.cuda.CudaPlatform" if is_cuda else None
builtin_platform_plugins = {
'cuda': cuda_platform_plugin,
}
def resolve_current_platform_cls_qualname() -> str:
# TODO(will): if we need to support other platforms, we should consider if
# vLLM's plugin architecture is suitable for our needs.
platform_cls_qualname = builtin_platform_plugins['cuda']()
return platform_cls_qualname
_current_platform = None
_init_trace: str = ''
if TYPE_CHECKING:
current_platform: Platform
def __getattr__(name: str):
if name == 'current_platform':
# lazy init current_platform.
# 1. out-of-tree platform plugins need `from vllm.platforms import
# Platform` so that they can inherit `Platform` class. Therefore,
# we cannot resolve `current_platform` during the import of
# `vllm.platforms`.
# 2. when users use out-of-tree platform plugins, they might run
# `import vllm`, some vllm internal code might access
# `current_platform` during the import, and we need to make sure
# `current_platform` is only resolved after the plugins are loaded
# (we have tests for this, if any developer violate this, they will
# see the test failures).
global _current_platform
if _current_platform is None:
platform_cls_qualname = resolve_current_platform_cls_qualname()
_current_platform = resolve_obj_by_qualname(
platform_cls_qualname)()
global _init_trace
_init_trace = "".join(traceback.format_stack())
return _current_platform
elif name in globals():
return globals()[name]
else:
raise AttributeError(
f"No attribute named '{name}' exists in {__name__}.")
__all__ = [
'Platform', 'PlatformEnum', 'current_platform',
"_init_trace"
]
-273
View File
@@ -1,273 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Code inside this file can safely assume cuda platform, e.g. importing
pynvml. However, it should not initialize cuda context.
"""
import os
from functools import lru_cache, wraps
from typing import (TYPE_CHECKING, Callable, List, Optional, Tuple, TypeVar,
Union)
import torch
from typing_extensions import ParamSpec
# import custom ops, trigger op registration
import fastvideo.v1.envs as envs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import import_pynvml
from .interface import DeviceCapability, Platform, PlatformEnum, _Backend
logger = init_logger(__name__)
_P = ParamSpec("_P")
_R = TypeVar("_R")
pynvml = import_pynvml()
# pytorch 2.5 uses cudnn sdpa by default, which will cause crash on some models
# see https://github.com/huggingface/diffusers/issues/9704 for details
torch.backends.cuda.enable_cudnn_sdp(False)
def device_id_to_physical_device_id(device_id: int) -> int:
if "CUDA_VISIBLE_DEVICES" in os.environ:
device_ids = os.environ["CUDA_VISIBLE_DEVICES"].split(",")
if device_ids == [""]:
msg = (
"CUDA_VISIBLE_DEVICES is set to empty string, which means"
" GPU support is disabled. If you are using ray, please unset"
" the environment variable `CUDA_VISIBLE_DEVICES` inside the"
" worker/actor. "
"Check https://github.com/vllm-project/vllm/issues/8402 for"
" more information.")
raise RuntimeError(msg)
physical_device_id = device_ids[device_id]
return int(physical_device_id)
else:
return device_id
def with_nvml_context(fn: Callable[_P, _R]) -> Callable[_P, _R]:
@wraps(fn)
def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _R:
pynvml.nvmlInit()
try:
return fn(*args, **kwargs)
finally:
pynvml.nvmlShutdown()
return wrapper
class CudaPlatformBase(Platform):
_enum = PlatformEnum.CUDA
device_name: str = "cuda"
device_type: str = "cuda"
dispatch_key: str = "CUDA"
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
@classmethod
def get_device_capability(cls,
device_id: int = 0
) -> Optional[DeviceCapability]:
raise NotImplementedError
@classmethod
def get_device_name(cls, device_id: int = 0) -> str:
raise NotImplementedError
@classmethod
def get_device_total_memory(cls, device_id: int = 0) -> int:
raise NotImplementedError
@classmethod
def is_async_output_supported(cls, enforce_eager: Optional[bool]) -> bool:
if enforce_eager:
logger.warning(
"To see benefits of async output processing, enable CUDA "
"graph. Since, enforce-eager is enabled, async output "
"processor cannot be used")
return False
return True
@classmethod
def is_full_nvlink(cls, device_ids: List[int]) -> bool:
raise NotImplementedError
@classmethod
def log_warnings(cls):
pass
@classmethod
def get_current_memory_usage(cls,
device: Optional[torch.types.Device] = None
) -> float:
torch.cuda.reset_peak_memory_stats(device)
return torch.cuda.max_memory_allocated(device)
@classmethod
def get_attn_backend_cls(cls, selected_backend, head_size, dtype,
kv_cache_dtype, block_size, use_v1,
use_mla) -> str:
raise NotImplementedError
@classmethod
def get_device_communicator_cls(cls) -> str:
return "fastvideo.v1.distributed.device_communicators.cuda_communicator.CudaCommunicator" # noqa
# NVML utils
# Note that NVML is not affected by `CUDA_VISIBLE_DEVICES`,
# all the related functions work on real physical device ids.
# the major benefit of using NVML is that it will not initialize CUDA
class NvmlCudaPlatform(CudaPlatformBase):
@classmethod
@lru_cache(maxsize=8)
@with_nvml_context
def get_device_capability(cls,
device_id: int = 0
) -> Optional[DeviceCapability]:
try:
physical_device_id = device_id_to_physical_device_id(device_id)
handle = pynvml.nvmlDeviceGetHandleByIndex(physical_device_id)
major, minor = pynvml.nvmlDeviceGetCudaComputeCapability(handle)
return DeviceCapability(major=major, minor=minor)
except RuntimeError:
return None
@classmethod
@lru_cache(maxsize=8)
@with_nvml_context
def has_device_capability(
cls,
capability: Union[Tuple[int, int], int],
device_id: int = 0,
) -> bool:
try:
return super().has_device_capability(capability, device_id)
except RuntimeError:
return False
@classmethod
@lru_cache(maxsize=8)
@with_nvml_context
def get_device_name(cls, device_id: int = 0) -> str:
physical_device_id = device_id_to_physical_device_id(device_id)
return cls._get_physical_device_name(physical_device_id)
@classmethod
@lru_cache(maxsize=8)
@with_nvml_context
def get_device_uuid(cls, device_id: int = 0) -> str:
physical_device_id = device_id_to_physical_device_id(device_id)
handle = pynvml.nvmlDeviceGetHandleByIndex(physical_device_id)
return pynvml.nvmlDeviceGetUUID(handle)
@classmethod
@lru_cache(maxsize=8)
@with_nvml_context
def get_device_total_memory(cls, device_id: int = 0) -> int:
physical_device_id = device_id_to_physical_device_id(device_id)
handle = pynvml.nvmlDeviceGetHandleByIndex(physical_device_id)
return int(pynvml.nvmlDeviceGetMemoryInfo(handle).total)
@classmethod
@with_nvml_context
def is_full_nvlink(cls, physical_device_ids: List[int]) -> bool:
"""
query if the set of gpus are fully connected by nvlink (1 hop)
"""
handles = [
pynvml.nvmlDeviceGetHandleByIndex(i) for i in physical_device_ids
]
for i, handle in enumerate(handles):
for j, peer_handle in enumerate(handles):
if i < j:
try:
p2p_status = pynvml.nvmlDeviceGetP2PStatus(
handle,
peer_handle,
pynvml.NVML_P2P_CAPS_INDEX_NVLINK,
)
if p2p_status != pynvml.NVML_P2P_STATUS_OK:
return False
except pynvml.NVMLError:
logger.exception(
"NVLink detection failed. This is normal if"
" your machine has no NVLink equipped.")
return False
return True
@classmethod
def _get_physical_device_name(cls, device_id: int = 0) -> str:
handle = pynvml.nvmlDeviceGetHandleByIndex(device_id)
return pynvml.nvmlDeviceGetName(handle)
@classmethod
@with_nvml_context
def log_warnings(cls):
device_ids: int = pynvml.nvmlDeviceGetCount()
if device_ids > 1:
device_names = [
cls._get_physical_device_name(i) for i in range(device_ids)
]
if (len(set(device_names)) > 1
and os.environ.get("CUDA_DEVICE_ORDER") != "PCI_BUS_ID"):
logger.warning(
"Detected different devices in the system: %s. Please"
" make sure to set `CUDA_DEVICE_ORDER=PCI_BUS_ID` to "
"avoid unexpected behavior.",
", ".join(device_names),
)
class NonNvmlCudaPlatform(CudaPlatformBase):
@classmethod
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability:
major, minor = torch.cuda.get_device_capability(device_id)
return DeviceCapability(major=major, minor=minor)
@classmethod
def get_device_name(cls, device_id: int = 0) -> str:
return torch.cuda.get_device_name(device_id)
@classmethod
def get_device_total_memory(cls, device_id: int = 0) -> int:
device_props = torch.cuda.get_device_properties(device_id)
return device_props.total_memory
@classmethod
def is_full_nvlink(cls, physical_device_ids: List[int]) -> bool:
logger.exception(
"NVLink detection not possible, as context support was"
" not found. Assuming no NVLink available.")
return False
# Autodetect either NVML-enabled or non-NVML platform
# based on whether NVML is available.
nvml_available = False
try:
try:
pynvml.nvmlInit()
nvml_available = True
except Exception:
# On Jetson, NVML is not supported.
nvml_available = False
finally:
if nvml_available:
pynvml.nvmlShutdown()
CudaPlatform = NvmlCudaPlatform if nvml_available else NonNvmlCudaPlatform
try:
from sphinx.ext.autodoc.mock import _MockModule
if not isinstance(pynvml, _MockModule):
CudaPlatform.log_warnings()
except ModuleNotFoundError:
CudaPlatform.log_warnings()
-194
View File
@@ -1,194 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import enum
import platform
import random
from platform import uname
from typing import TYPE_CHECKING, NamedTuple, Optional, Tuple, Union
import numpy as np
import torch
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class _Backend(enum.Enum):
FLASH_ATTN = enum.auto()
TORCH_SDPA = enum.auto()
NO_ATTENTION = enum.auto()
class PlatformEnum(enum.Enum):
CUDA = enum.auto()
OOT = enum.auto()
UNSPECIFIED = enum.auto()
class DeviceCapability(NamedTuple):
major: int
minor: int
def as_version_str(self) -> str:
return f"{self.major}.{self.minor}"
def to_int(self) -> int:
"""
Express device capability as an integer ``<major><minor>``.
It is assumed that the minor version is always a single digit.
"""
assert 0 <= self.minor < 10
return self.major * 10 + self.minor
class Platform:
_enum: PlatformEnum
device_name: str
device_type: str
# available dispatch keys:
# check https://github.com/pytorch/pytorch/blob/313dac6c1ca0fa0cde32477509cce32089f8532a/torchgen/model.py#L134 # noqa
# use "CPU" as a fallback for platforms not registered in PyTorch
dispatch_key: str = "CPU"
supported_quantization: list[str] = []
def is_cuda(self) -> bool:
return self._enum == PlatformEnum.CUDA
def is_out_of_tree(self) -> bool:
return self._enum == PlatformEnum.OOT
def is_cuda_alike(self) -> bool:
"""Stateless version of :func:`torch.cuda.is_available`."""
return self._enum in (PlatformEnum.CUDA, PlatformEnum.ROCM)
@classmethod
def get_attn_backend_cls(cls, selected_backend: _Backend, head_size: int,
dtype: torch.dtype, kv_cache_dtype: Optional[str],
block_size: int, use_v1: bool,
use_mla: bool) -> str:
"""Get the attention backend class of a device."""
return ""
@classmethod
def get_device_capability(
cls,
device_id: int = 0,
) -> Optional[DeviceCapability]:
"""Stateless version of :func:`torch.cuda.get_device_capability`."""
return None
@classmethod
def has_device_capability(
cls,
capability: Union[Tuple[int, int], int],
device_id: int = 0,
) -> bool:
"""
Test whether this platform is compatible with a device capability.
The ``capability`` argument can either be:
- A tuple ``(major, minor)``.
- An integer ``<major><minor>``. (See :meth:`DeviceCapability.to_int`)
"""
current_capability = cls.get_device_capability(device_id=device_id)
if current_capability is None:
return False
if isinstance(capability, tuple):
return current_capability >= capability
return current_capability.to_int() >= capability
@classmethod
def get_device_name(cls, device_id: int = 0) -> str:
"""Get the name of a device."""
raise NotImplementedError
@classmethod
def get_device_uuid(cls, device_id: int = 0) -> str:
"""Get the uuid of a device, e.g. the PCI bus ID."""
raise NotImplementedError
@classmethod
def get_device_total_memory(cls, device_id: int = 0) -> int:
"""Get the total memory of a device in bytes."""
raise NotImplementedError
@classmethod
def is_async_output_supported(cls, enforce_eager: Optional[bool]) -> bool:
"""
Check if the current platform supports async output.
"""
raise NotImplementedError
@classmethod
def inference_mode(cls):
"""A device-specific wrapper of `torch.inference_mode`.
This wrapper is recommended because some hardware backends such as TPU
do not support `torch.inference_mode`. In such a case, they will fall
back to `torch.no_grad` by overriding this method.
"""
return torch.inference_mode(mode=True)
@classmethod
def seed_everything(cls, seed: Optional[int] = None) -> None:
"""
Set the seed of each random module.
`torch.manual_seed` will set seed on all devices.
Loosely based on: https://github.com/Lightning-AI/pytorch-lightning/blob/2.4.0/src/lightning/fabric/utilities/seed.py#L20
"""
if seed is not None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
@classmethod
def verify_model_arch(cls, model_arch: str) -> None:
"""
Verify whether the current platform supports the specified model
architecture.
- This will raise an Error or Warning based on the model support on
the current platform.
- By default all models are considered supported.
"""
pass
@classmethod
def verify_quantization(cls, quant: str) -> None:
"""
Verify whether the quantization is supported by the current platform.
"""
if cls.supported_quantization and \
quant not in cls.supported_quantization:
raise ValueError(
f"{quant} quantization is currently not supported in "
f"{cls.device_name}.")
@classmethod
def get_current_memory_usage(cls,
device: Optional[torch.types.Device] = None
) -> float:
"""
Return the memory usage in bytes.
"""
raise NotImplementedError
@classmethod
def get_device_communicator_cls(cls) -> str:
"""
Get device specific communicator class for distributed communication.
"""
return "fastvideo.v1.distributed.device_communicators.base_device_communicator.DeviceCommunicatorBase" # noqa
class UnspecifiedPlatform(Platform):
_enum = PlatformEnum.UNSPECIFIED
device_type = ""
@@ -1,206 +0,0 @@
import os
import imageio
import numpy as np
import torch
import torchvision
from einops import rearrange
import sys
# Fix the import path
from fastvideo.v1.inference_engine import InferenceEngine
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.inference_args import prepare_inference_args
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
from fastvideo.v1.logger import logger
from fastvideo.v1.models.loader.fsdp_load import get_param_names_mapping
from safetensors.torch import safe_open, save_file
import shutil
def initialize_distributed_and_parallelism(inference_args: InferenceArgs):
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(local_rank)
init_distributed_environment(
world_size=world_size,
rank=rank,
local_rank=local_rank
)
device_str = f"cuda:{local_rank}"
inference_args.device_str = device_str
inference_args.device = torch.device(device_str)
initialize_model_parallel(
sequence_model_parallel_size=inference_args.sp_size,
tensor_model_parallel_size=inference_args.tp_size,
)
def convert_and_save_lora_weights(input_folder, output_folder, lora_param_mapping_fn):
# Check if the input folder exists
if not os.path.exists(input_folder):
raise ValueError(f"Input folder {input_folder} not found")
# Create output folder if it doesn't exist
os.makedirs(output_folder, exist_ok=True)
# Process safetensors file
safetensors_file = os.path.join(input_folder, "pytorch_lora_weights.safetensors")
if os.path.exists(safetensors_file):
# Load all tensors from the safetensors file
tensors = {}
with safe_open(safetensors_file, framework="pt", device="cpu") as f:
tensor_names = f.keys()
print(f"Found {len(tensor_names)} tensors in {safetensors_file}")
for name in tensor_names:
tensors[name] = f.get_tensor(name)
# Extract base names (without LoRA suffix) and LoRA suffixes
base_names_dict = {}
lora_suffixes = {}
for name in tensor_names:
if name.endswith(".lora_A.weight") or name.endswith(".lora_B.weight"):
for suffix in [".lora_A.weight", ".lora_B.weight"]:
if name.endswith(suffix):
base_name = name[:-len(suffix)] + ".weight"
base_names_dict[name] = base_name
lora_suffixes[name] = suffix
break
else:
base_names_dict[name] = name
lora_suffixes[name] = ""
# Apply the mapping function to get all mappings at once
name_mappings = {}
for full_name, base_name in base_names_dict.items():
try:
# First remove "transformer." prefix if it exists
has_transformer_prefix = False
processed_base_name = base_name
if base_name.startswith("transformer."):
processed_base_name = base_name[12:] # Remove "transformer." (12 characters)
has_transformer_prefix = True
# Apply the mapping function to the processed name
target_base_name, merge_index, total_splitted_params = lora_param_mapping_fn(processed_base_name)
# Add back the "transformer." prefix if it was removed
if target_base_name is not None and has_transformer_prefix:
target_base_name = "transformer." + target_base_name
if target_base_name is not None:
name_mappings[full_name] = (target_base_name, merge_index, total_splitted_params)
except Exception as e:
print(f"Error mapping parameter {base_name}: {e}")
# Create the converted weights using the mappings
converted_weights = {}
for name, tensor in tensors.items():
suffix = lora_suffixes.get(name, "")
if name in name_mappings:
target_base_name, merge_index, _ = name_mappings[name]
# Reconstruct full parameter name with LoRA suffix
if target_base_name.endswith('.weight'):
target_base_name_without_suffix = target_base_name[:-7] # 移除 .weight
new_key = target_base_name_without_suffix + suffix
else:
new_key = target_base_name + suffix
if merge_index is not None:
new_key += f"_{merge_index}"
converted_weights[new_key] = tensor
else:
# Keep original name if no mapping was found
converted_weights[name] = tensor
print(f"Converted {len(converted_weights)} parameters")
# Save the converted weights
output_safetensors_file = os.path.join(output_folder, "pytorch_lora_weights.safetensors")
save_file(converted_weights, output_safetensors_file)
print(f"Converted weights saved to {output_safetensors_file}")
else:
print(f"Warning: No safetensors file found in {input_folder}")
# Copy other files (like config and optimizer state) without modification
for filename in os.listdir(input_folder):
if filename != "pytorch_lora_weights.safetensors":
input_path = os.path.join(input_folder, filename)
output_path = os.path.join(output_folder, filename)
if os.path.isfile(input_path):
shutil.copy2(input_path, output_path)
print(f"Copied {filename} to output folder")
print(f"LoRA conversion complete: {input_folder} → {output_folder}")
def main(inference_args: InferenceArgs):
initialize_distributed_and_parallelism(inference_args)
engine = InferenceEngine.create_engine(
inference_args,
)
if inference_args.prompt_path is not None:
with open(inference_args.prompt_path) as f:
prompts = [line.strip() for line in f.readlines()]
else:
prompts = [inference_args.prompt]
# from IPython import embed; embed()
# # convert lora format
# input_path = "data/Hunyuan-Black-Myth-Wukong-lora-weight/"
# output_path = "data/Hunyuan-Black-Myth-Wukong-lora-weight_converted"
# # from IPython import embed; embed()
# lora_param_mapping_fn = get_param_names_mapping(engine.pipeline.transformer._param_names_mapping)
# converted_weights = convert_and_save_lora_weights(
# input_path,
# output_path,
# lora_param_mapping_fn,
# )
# Process each prompt
for prompt in prompts:
# lora_checkpoint = "data/Hunyuan-Black-Myth-Wukong-lora-weight_converted/"
# if lora_checkpoint:
# import json
# print(f"Loading LoRA weights from lora_checkpoint: {lora_checkpoint}")
# config_path = os.path.join(lora_checkpoint, "lora_config.json")
# with open(config_path, "r") as f:
# lora_config_dict = json.load(f)
# rank = lora_config_dict["lora_params"]["lora_rank"]
# lora_alpha = lora_config_dict["lora_params"]["lora_alpha"]
# lora_scaling = lora_alpha / rank
# engine.pipeline.load_lora_weights(lora_checkpoint, adapter_name="default")
# from IPython import embed; embed()
# engine.pipeline.set_adapters(["default"], [lora_scaling])
# print(f"Successfully Loaded LoRA weights from {lora_checkpoint}")
outputs = engine.run(
prompt=prompt,
inference_args=inference_args,
)
# img_attn_qkv, img_attn_proj, linear1, linear2
# Process outputs
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
# Save video
os.makedirs(os.path.dirname(inference_args.output_path), exist_ok=True)
imageio.mimsave(
os.path.join(inference_args.output_path, f"{prompt[:100]}.mp4"),
frames,
fps=inference_args.fps
)
if __name__ == "__main__":
inference_args = prepare_inference_args(sys.argv[1:])
main(inference_args)
-18
View File
@@ -1,18 +0,0 @@
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
import os
def main():
print(os.environ["RANK"])
print(os.environ["WORLD_SIZE"])
print(os.environ["LOCAL_RANK"])
print(os.environ["MASTER_ADDR"])
print(os.environ["MASTER_PORT"])
print(os.environ["RANK"])
print(os.environ["WORLD_SIZE"])
print(os.environ["LOCAL_RANK"])
print(os.environ["MASTER_ADDR"])
if __name__ == "__main__":
main()
-19
View File
@@ -1,19 +0,0 @@
#!/bin/bash
num_gpus=8
torchrun --standalone --nnodes=1 --nproc_per_node=$num_gpus \
--master_port 29503 \
tp_example.py
num_gpus=2
torchrun --standalone --nnodes=4 --nproc_per_node=$num_gpus \
--master_port 29503 \
test_hunyuanvideo.py --sequence_model_parallel_size $num_gpus
num_gpus=2
torchrun --standalone --nnodes=1 --nproc_per_node=$num_gpus \
--master_port 29503 \
fastvideo/v1/tests/test_hunyuanvideo_load.py --sequence_model_parallel_size $num_gpus
-207
View File
@@ -1,207 +0,0 @@
from fastvideo.v1.models.encoders.clip import CLIPTextModel
# TODO: check if correct
from fastvideo.models.hunyuan.text_encoder import TextEncoder, load_text_encoder, load_tokenizer
import os
import torch
import torch.nn as nn
import argparse
import numpy as np
from fastvideo.v1.logger import init_logger
from transformers import AutoConfig
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
# from fastvideo.v1.models.hunyuan.text_encoder import load_text_encoder, load_tokenizer
logger = init_logger(__name__)
def initialize_identical_weights(model1, model2, seed=42):
"""Initialize both models with identical weights using a fixed seed for reproducibility."""
# Get all parameters from both models
params1 = dict(model1.named_parameters())
params2 = dict(model2.named_parameters())
# Initialize each layer with identical values
with torch.no_grad():
# Initialize weights
for name1, param1 in params1.items():
if 'weight' in name1:
# Set seed before each weight initialization
torch.manual_seed(seed)
nn.init.normal_(param1, mean=0.0, std=0.05)
for name2, param2 in params2.items():
if 'weight' in name2:
# Reset seed to get same initialization
torch.manual_seed(seed)
nn.init.normal_(param2, mean=0.0, std=0.05)
# Initialize biases
for name1, param1 in params1.items():
if 'bias' in name1:
torch.manual_seed(seed)
nn.init.normal_(param1, mean=0.0, std=0.05)
param1.data = param1.data.to(torch.bfloat16)
for name2, param2 in params2.items():
if 'bias' in name2:
torch.manual_seed(seed)
nn.init.normal_(param2, mean=0.0, std=0.05)
param2.data = param2.data.to(torch.bfloat16)
logger.info("Both models initialized with identical weights in bfloat16")
return model1, model2
def setup_args():
parser = argparse.ArgumentParser(description='CLIP Text Encoder Test')
parser.add_argument('--model-path', type=str, default="openai/clip-vit-large-patch14",
help='Path to the CLIP model')
parser.add_argument('--precision', type=str, default="float16",
help='Precision to use for the model (float32, float16, bfloat16)')
return parser.parse_args()
def test_clip_encoder():
init_distributed_environment(world_size=1, rank=0, distributed_init_method="env://", local_rank=0, backend="nccl")
initialize_model_parallel(tensor_model_parallel_size=1, sequence_model_parallel_size=1, backend="nccl")
args = setup_args()
# Set fixed random seed for reproducibility
torch.manual_seed(42)
np.random.seed(42)
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
# Initialize the two model implementations
logger.info(f"Loading models from {args.model_path}")
model_path = "data/hunyuanvideo-community/HunyuanVideo/text_encoder_2"
# config = json.load(open(os.path.join(model_path, "config.json")))
hf_config = AutoConfig.from_pretrained(model_path)
print(hf_config)
print(hf_config.use_return_dict)
# Load our implementation using the loader from text_encoder/__init__.py
model1, _ = load_text_encoder(
text_encoder_type="clipL",
text_encoder_precision='fp16',
text_encoder_path=model_path,
logger=logger,
device=device
)
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
loader = TextEncoderLoader()
args.device_str = "cuda:0"
model2 = loader.load_model(model_path, hf_config, device)
# Load the HuggingFace implementation directly
# model2 = CLIPTextModel(hf_config)
model2 = model2.to(torch.float16)
model2 = model2.to(device)
model2.eval()
# Sanity check weights between the two models
logger.info("Comparing model weights for sanity check...")
params1 = dict(model1.named_parameters())
params2 = dict(model2.named_parameters())
# Check number of parameters
logger.info(f"Model1 has {len(params1)} parameters")
logger.info(f"Model2 has {len(params2)} parameters")
# Compare a few key parameters
# weight_diffs = []
# for (name1, param1), (name2, param2) in zip(
# sorted(params1.items()), sorted(params2.items())
# ):
# # if len(weight_diffs) < 5: # Just check a few parameters
# max_diff = torch.max(torch.abs(param1 - param2)).item()
# mean_diff = torch.mean(torch.abs(param1 - param2)).item()
# weight_diffs.append((name1, name2, max_diff, mean_diff))
# logger.info(f"Parameter: {name1} vs {name2}")
# logger.info(f" Max diff: {max_diff}, Mean diff: {mean_diff}")
# Load tokenizer
tokenizer, _ = load_tokenizer(
tokenizer_type="clipL",
tokenizer_path=args.model_path,
logger=logger
)
# Test with some sample prompts
prompts = [
"a photo of a cat",
"a beautiful landscape with mountains",
"an astronaut riding a horse on the moon"
]
logger.info("Testing CLIP text encoder with sample prompts")
with torch.no_grad():
for prompt in prompts:
logger.info(f"Testing prompt: '{prompt}'")
# Tokenize the prompt
tokens = tokenizer(
prompt,
padding="max_length",
max_length=77,
truncation=True,
return_tensors="pt"
).to(device)
# Get embeddings from our implementation
outputs1 = model1(
input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
output_hidden_states=True
)
logger.info(f"Testing model2")
# Get embeddings from HuggingFace implementation
outputs2 = model2(
input_ids=tokens.input_ids,
# attention_mask=tokens.attention_mask,
output_hidden_states=True
)
# Compare last hidden states
last_hidden_state1 = outputs1.last_hidden_state[tokens.attention_mask==1]
last_hidden_state2 = outputs2.last_hidden_state[tokens.attention_mask==1]
# print("last_hidden_state1", last_hidden_state1)
# print("last_hidden_state2", last_hidden_state2)
assert last_hidden_state1.shape == last_hidden_state2.shape, \
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
max_diff_hidden = torch.max(torch.abs(last_hidden_state1 - last_hidden_state2))
mean_diff_hidden = torch.mean(torch.abs(last_hidden_state1 - last_hidden_state2))
logger.info(f"Maximum difference in last hidden states: {max_diff_hidden.item()}")
logger.info(f"Mean difference in last hidden states: {mean_diff_hidden.item()}")
# Compare pooler outputs
pooler_output1 = outputs1.pooler_output
pooler_output2 = outputs2.pooler_output
assert pooler_output1.shape == pooler_output2.shape, \
f"Pooler output shapes don't match: {pooler_output1.shape} vs {pooler_output2.shape}"
max_diff_pooler = torch.max(torch.abs(pooler_output1 - pooler_output2))
mean_diff_pooler = torch.mean(torch.abs(pooler_output1 - pooler_output2))
logger.info(f"Maximum difference in pooler outputs: {max_diff_pooler.item()}")
logger.info(f"Mean difference in pooler outputs: {mean_diff_pooler.item()}")
# Check if outputs are similar (allowing for small numerical differences)
assert max_diff_hidden < 1e-4, \
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
assert max_diff_pooler < 1e-4, \
f"Pooler outputs differ significantly: max diff = {max_diff_pooler.item()}"
logger.info("Test passed! Both CLIP text encoder implementations produce similar outputs.")
logger.info("Test completed successfully")
if __name__ == "__main__":
test_clip_encoder()
-161
View File
@@ -1,161 +0,0 @@
from fastvideo.v1.models.vaes.hunyuanvae import AutoencoderKLHunyuanVideo as MyHunyuanVAE
from diffusers import AutoencoderKLHunyuanVideo as DiffusersHunyuanVAE
import os
import torch
import torch.nn as nn
import argparse
import numpy as np
import json
from fastvideo.v1.logger import init_logger
from safetensors.torch import load_file
logger = init_logger(__name__)
def initialize_identical_weights(model1, model2, seed=42):
"""Initialize both models with identical weights using a fixed seed for reproducibility."""
# Get all parameters from both models
params1 = dict(model1.named_parameters())
params2 = dict(model2.named_parameters())
# Initialize each layer with identical values
with torch.no_grad():
# Initialize weights
for name1, param1 in params1.items():
if 'weight' in name1:
# Set seed before each weight initialization
torch.manual_seed(seed)
nn.init.normal_(param1, mean=0.0, std=0.05)
for name2, param2 in params2.items():
if 'weight' in name2:
# Reset seed to get same initialization
torch.manual_seed(seed)
nn.init.normal_(param2, mean=0.0, std=0.05)
# Initialize biases
for name1, param1 in params1.items():
if 'bias' in name1:
torch.manual_seed(seed)
nn.init.normal_(param1, mean=0.0, std=0.05)
param1.data = param1.data.to(torch.bfloat16)
for name2, param2 in params2.items():
if 'bias' in name2:
torch.manual_seed(seed)
nn.init.normal_(param2, mean=0.0, std=0.05)
param2.data = param2.data.to(torch.bfloat16)
logger.info("Both models initialized with identical weights in bfloat16")
return model1, model2
def setup_args():
parser = argparse.ArgumentParser(description='HunyuanVAE Test')
parser.add_argument('--in-channels', type=int, default=4,
help='Number of input channels')
parser.add_argument('--out-channels', type=int, default=4,
help='Number of output channels')
parser.add_argument('--latent-channels', type=int, default=4,
help='Number of latent channels')
return parser.parse_args()
def test_hunyuan_vae():
args = setup_args()
# Set fixed random seed for reproducibility
torch.manual_seed(42)
np.random.seed(42)
# Model parameters
in_channels = args.in_channels
out_channels = args.out_channels
latent_channels = args.latent_channels
device = torch.device("cuda:0")
print
# Initialize the two model implementations
path = "data/hunyuanvideo-community/HunyuanVideo/vae"
config_path = os.path.join(path, "config.json")
config = json.load(open(config_path))
config.pop("_class_name")
config.pop("_diffusers_version")
model1 = MyHunyuanVAE(
**config
).to(torch.bfloat16)
model2 = DiffusersHunyuanVAE(
**config
).to(torch.bfloat16)
loaded = load_file(os.path.join(path, "diffusion_pytorch_model.safetensors"))
model1.load_state_dict(loaded)
model2.load_state_dict(loaded)
# Set both models to eval mode
model1.eval()
model2.eval()
# Move to GPU
model1 = model1.to(device)
model2 = model2.to(device)
model1.enable_tiling(
tile_sample_min_height=32,
tile_sample_min_width=32,
tile_sample_min_num_frames=8,
tile_sample_stride_height=16,
tile_sample_stride_width=16,
tile_sample_stride_num_frames=4
)
model2.enable_tiling(
tile_sample_min_height=32,
tile_sample_min_width=32,
tile_sample_min_num_frames=8,
tile_sample_stride_height=16,
tile_sample_stride_width=16,
tile_sample_stride_num_frames=4
)
# Create identical inputs for both models
batch_size = 1
# Video input [B, C, T, H, W]
input_tensor = torch.randn(batch_size, 3, 21, 64, 64, device=device, dtype=torch.bfloat16)
# Disable gradients for inference
with torch.no_grad():
# Test encoding
logger.info("Testing encoding...")
latent1 = model1.encode(input_tensor).mean
print("--------------------------------")
latent2 = model2.encode(input_tensor).latent_dist.mean
# Check if latents have the same shape
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
# Check if latents are similar
max_diff_encode = torch.max(torch.abs(latent1 - latent2))
mean_diff_encode = torch.mean(torch.abs(latent1 - latent2))
logger.info(f"Maximum difference between encoded latents: {max_diff_encode.item()}")
logger.info(f"Mean difference between encoded latents: {mean_diff_encode.item()}")
assert max_diff_encode < 1e-4, f"Encoded latents differ significantly: max diff = {max_diff_encode.item()}"
# Test decoding
logger.info("Testing decoding...")
output1 = model1.decode(latent1)
output2 = model2.decode(latent2).sample
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
# Check if outputs are similar
max_diff_decode = torch.max(torch.abs(output1 - output2))
mean_diff_decode = torch.mean(torch.abs(output1 - output2))
logger.info(f"Maximum difference between decoded outputs: {max_diff_decode.item()}")
logger.info(f"Mean difference between decoded outputs: {mean_diff_decode.item()}")
assert max_diff_decode < 1e-4, f"Decoded outputs differ significantly: max diff = {max_diff_decode.item()}"
logger.info("Test passed! Both VAE implementations produce similar outputs.")
logger.info("Test completed successfully")
if __name__ == "__main__":
test_hunyuan_vae()
-233
View File
@@ -1,233 +0,0 @@
from itertools import chain
import os
import torch
import torch.nn as nn
import argparse
import numpy as np
from fastvideo.v1.logger import init_logger
from fastvideo.v1.distributed.parallel_state import (
init_distributed_environment,
initialize_model_parallel,
get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size,
destroy_model_parallel,
destroy_distributed_environment,
cleanup_dist_env_and_memory
)
from fastvideo.v1.utils.parallel_states import initialize_sequence_parallel_state
from torch.distributed._composable.fsdp import CPUOffloadPolicy, fully_shard
from torch.distributed.device_mesh import init_device_mesh
from fastvideo.v1.models.loader.fsdp_load import shard_model
from fastvideo.v1.models.dits.hunyuanvideo import HunyuanVideoTransformer3DModel as HunyuanVideoDit
from fastvideo.v1.models.hunyuan.modules.models import HYVideoDiffusionTransformer
from fastvideo.v1.models.hunyuan_hf.modeling_hunyuan import HunyuanVideoTransformer3DModel
logger = init_logger(__name__)
def initialize_identical_weights(model1, model2, seed=42):
"""Initialize both models with identical weights using a fixed seed for reproducibility."""
# Get all parameters from both models
params1 = dict(model1.named_parameters())
params2 = dict(model2.named_parameters())
# Initialize each layer with identical values
with torch.no_grad():
# Initialize weights
for name1, param1 in params1.items():
if 'weight' in name1:
# Set seed before each weight initialization
torch.manual_seed(seed)
nn.init.normal_(param1, mean=0.0, std=0.05)
for name2, param2 in params2.items():
if 'weight' in name2:
# Reset seed to get same initialization
torch.manual_seed(seed)
nn.init.normal_(param2, mean=0.0, std=0.05)
# Initialize biases
for name1, param1 in params1.items():
if 'bias' in name1:
torch.manual_seed(seed)
nn.init.normal_(param1, mean=0.0, std=0.05)
param1.data = param1.data.to(torch.bfloat16)
for name2, param2 in params2.items():
if 'bias' in name2:
torch.manual_seed(seed)
nn.init.normal_(param2, mean=0.0, std=0.05)
param2.data = param2.data.to(torch.bfloat16)
logger.info("Both models initialized with identical weights in bfloat16")
return model1, model2
def setup_args():
parser = argparse.ArgumentParser(description='Distributed HunyuanVideo Test')
parser.add_argument('--sequence_model_parallel_size', type=int, default=1,
help='Degree of sequence model parallelism')
parser.add_argument('--hidden-size', type=int, default=128,
help='Hidden size for the model')
parser.add_argument('--heads-num', type=int, default=4,
help='Number of attention heads')
parser.add_argument('--double-blocks-depth', type=int, default=2,
help='Number of double stream blocks')
parser.add_argument('--single-blocks-depth', type=int, default=2,
help='Number of single stream blocks')
return parser.parse_args()
def test_hunyuanvideo_distributed():
args = setup_args()
# Initialize distributed environment
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))
logger.info(f"Initializing process: rank={rank}, local_rank={local_rank}, world_size={world_size}")
# Initialize distributed environment
init_distributed_environment(
world_size=world_size,
rank=rank,
local_rank=local_rank
)
# Initialize tensor model parallel groups
initialize_model_parallel(
sequence_model_parallel_size=args.sequence_model_parallel_size
)
initialize_sequence_parallel_state(world_size)
# Get tensor parallel info
sp_rank = get_sequence_model_parallel_rank()
sp_world_size = get_sequence_model_parallel_world_size()
logger.info(f"Process rank {rank} initialized with SP rank {sp_rank} in SP world size {sp_world_size}")
# Set fixed random seed for reproducibility
torch.manual_seed(42)
np.random.seed(42)
# Small model parameters for testing
hidden_size = args.hidden_size
heads_num = args.heads_num
mm_double_blocks_depth = args.double_blocks_depth
mm_single_blocks_depth = args.single_blocks_depth
patch_size = [1, 2, 2]
torch.cuda.set_device(f"cuda:{local_rank}")
# Initialize the two model implementations
model1 = HunyuanVideoDit(
patch_size=2,
patch_size_t=1,
in_channels=4,
out_channels=4,
attention_head_dim=hidden_size//heads_num,
num_attention_heads=heads_num,
num_layers=mm_double_blocks_depth,
num_single_layers=mm_single_blocks_depth,
rope_axes_dim=[8, 16, 8], # sum = hidden_size // heads_num = 32
dtype=torch.bfloat16
).to(torch.bfloat16)
model2 = HYVideoDiffusionTransformer(
patch_size=patch_size,
in_channels=4,
hidden_size=hidden_size,
heads_num=heads_num,
mm_double_blocks_depth=mm_double_blocks_depth,
mm_single_blocks_depth=mm_single_blocks_depth,
rope_dim_list=[8, 16, 8], # sum = hidden_size // heads_num = 32
dtype=torch.bfloat16
).to(torch.bfloat16)
# print("--------------------------------")
# for name, param in model3.named_parameters():
# print(name)
# import pdb; pdb.set_trace()
# # Initialize with identical weights
model1, model2 = initialize_identical_weights(model1, model2, seed=42)
device_mesh = init_device_mesh(
"cuda",
mesh_shape=(sp_world_size,),
mesh_dim_names=("dp", ),
)
shard_model(model1, cpu_offload=False, reshard_after_forward=True)
for n, p in chain(model1.named_parameters(), model1.named_buffers()):
if p.is_meta:
raise RuntimeError(f"Unexpected param or buffer {n} on meta device.")
for p in model1.parameters():
p.requires_grad = False
# Set both models to eval mode
model1.eval()
model2.eval()
# Move to GPU based on local rank (0 or 1 for 2 GPUs)
device = torch.device(f"cuda:{local_rank}")
model1 = model1.to(device)
model2 = model2.to(device)
# Create identical inputs for both models
batch_size = 1
seq_len = 3
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size, 4, 8, 16, 16, device=device, dtype=torch.bfloat16)
chunk_per_rank = hidden_states.shape[2] // sp_world_size
hidden_states = hidden_states[:, :, sp_rank * chunk_per_rank:(sp_rank + 1) * chunk_per_rank]
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size, seq_len + 1, 4096, device=device, dtype=torch.bfloat16)
# Timestep
timestep = torch.tensor([500], device=device, dtype=torch.bfloat16)
# Attention mask for text
encoder_attention_mask = torch.ones(batch_size, seq_len, device=device, dtype=torch.bfloat16)
guidance = torch.tensor([1.0], device=device, dtype=torch.bfloat16)
# Disable gradients for inference
with torch.no_grad():
output1 = model1(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
)
print("--------------------------------")
output2, _ = model2(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
encoder_attention_mask=encoder_attention_mask,
)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
logger.info(f"Maximum difference between outputs: {max_diff.item()}")
# mean diff
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info(f"Mean difference between outputs: {mean_diff.item()}")
# diff sum
diff_sum = torch.sum(torch.abs(output1 - output2))
logger.info(f"Diff sum between outputs: {diff_sum.item()}")
# sum
sum_output1 = torch.sum(output1.float())
sum_output2 = torch.sum(output2.float())
logger.info(f"Rank {sp_rank} Sum of output1: {sum_output1.item()}")
logger.info(f"Rank {sp_rank} Sum of output2: {sum_output2.item()}")
# The outputs should be very close if not identical
assert max_diff < 1e-3, f"Outputs differ significantly: max diff = {max_diff.item()}" # Increased tolerance for bf16
logger.info("Test passed! Both model implementations produce the same outputs.")
# Clean up
logger.info("Cleaning up distributed environment")
destroy_model_parallel()
destroy_distributed_environment()
cleanup_dist_env_and_memory()
logger.info("Test completed successfully")
if __name__ == "__main__":
test_hunyuanvideo_distributed()
@@ -1,195 +0,0 @@
import os
import torch
import torch.nn as nn
import argparse
import numpy as np
from fastvideo.v1.logger import init_logger
from fastvideo.v1.distributed.parallel_state import (
init_distributed_environment,
initialize_model_parallel,
get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size,
destroy_model_parallel,
destroy_distributed_environment,
cleanup_dist_env_and_memory
)
import json
from fastvideo.v1.models.dits.hunyuanvideo import HunyuanVideoTransformer3DModel as HunyuanVideoDit
from fastvideo.models.hunyuan.modules.models import HUNYUAN_VIDEO_CONFIG
from fastvideo.models.hunyuan.modules.models import HYVideoDiffusionTransformer
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
from fastvideo.models.hunyuan.inference import Inference
import glob
logger = init_logger(__name__)
def setup_args():
parser = argparse.ArgumentParser(description='Distributed HunyuanVideo Test')
parser.add_argument('--sequence_model_parallel_size', type=int, default=1,
help='Degree of sequence model parallelism')
return parser.parse_args()
def test_hunyuanvideo_distributed():
args = setup_args()
# Initialize distributed environment
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))
logger.info(f"Initializing process: rank={rank}, local_rank={local_rank}, world_size={world_size}")
# Initialize distributed environment
init_distributed_environment(
world_size=world_size,
rank=rank,
local_rank=local_rank
)
torch.cuda.set_device(f"cuda:{local_rank}")
# Initialize tensor model parallel groups
initialize_model_parallel(
sequence_model_parallel_size=args.sequence_model_parallel_size
)
# Get tensor parallel info
sp_rank = get_sequence_model_parallel_rank()
sp_world_size = get_sequence_model_parallel_world_size()
logger.info(f"Process rank {rank} initialized with SP rank {sp_rank} in SP world size {sp_world_size}")
# load data/hunyuanvideo_community/transformer/config.json
with open("data/hunyuanvideo-community/HunyuanVideo/transformer/config.json", "r") as f:
config = json.load(f)
# remove "_class_name": "HunyuanVideoTransformer3DModel", "_diffusers_version": "0.32.0.dev0",
# TODO: write normalize config function
config.pop("_class_name")
config.pop("_diffusers_version")
# load data/hunyuanvideo_community/transformer/*.safetensors
weight_dir_list = glob.glob("data/hunyuanvideo-community/HunyuanVideo/transformer/*.safetensors")
# to str
weight_dir_list = [str(path) for path in weight_dir_list]
model1 = load_fsdp_model(
HunyuanVideoDit,
init_params=config,
weight_dir_list=weight_dir_list,
device=torch.device(f"cuda:{local_rank}"),
cpu_offload=False
)
# successfully sharded the model (hunyuanvideo bf16 should take around 26GB in total)
total_params = sum(p.numel() for p in model1.parameters())
logger.info(f"Total parameters: {total_params / 1e9}B")
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
model2 = HYVideoDiffusionTransformer(
in_channels=16,
out_channels=16,
**HUNYUAN_VIDEO_CONFIG["HYVideo-T/2-cfgdistill"],
device=torch.device(f"cuda:{local_rank}"),
dtype=torch.bfloat16
).bfloat16()
# data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt
state_dict = torch.load("/mbz/users/hao.zhang/peiyuan/FastVideo/data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt", map_location=lambda storage, loc: storage)["module"]
model2.load_state_dict(state_dict, strict=True)
model2.to(torch.device(f"cuda:{local_rank}")).bfloat16()
print("load state dict done")
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info(f"Model 2 weight sum: {weight_sum_model2}")
logger.info(f"Model 2 weight mean: {weight_mean_model2}")
# Set both models to eval mode
model1.eval()
model2.eval()
# Create random inputs for testing
batch_size = 1
seq_len = 3
device = torch.device(f"cuda:{local_rank}")
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size, 16, 8, 16, 16, device=device, dtype=torch.bfloat16)
chunk_per_rank = hidden_states.shape[2] // sp_world_size
hidden_states = hidden_states[:, :, sp_rank * chunk_per_rank:(sp_rank + 1) * chunk_per_rank]
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size, seq_len + 1, 4096, device=device, dtype=torch.bfloat16)
# Timestep
timestep = torch.tensor([500], device=device, dtype=torch.bfloat16)
# Attention mask for text
encoder_attention_mask = torch.ones(batch_size, seq_len, device=device, dtype=torch.bfloat16)
guidance = torch.tensor([1.0], device=device, dtype=torch.bfloat16)
# Disable gradients for inference
with torch.no_grad():
# Run inference on model1
logger.info(f"Running inference on model1")
output1 = model1(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
)
logger.info("Model 1 inference completed")
# Run inference on model2
output2, _ = model2(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
encoder_attention_mask=encoder_attention_mask,
)
logger.info("Model 2 inference completed")
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
# Compare weight sums and means
logger.info(f"Model 1 weight sum: {weight_sum_model1}")
logger.info(f"Model 2 weight sum: {weight_sum_model2}")
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info(f"Weight sum difference: {weight_sum_diff}")
logger.info(f"Model 1 weight mean: {weight_mean_model1}")
logger.info(f"Model 2 weight mean: {weight_mean_model2}")
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info(f"Weight mean difference: {weight_mean_diff}")
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
assert max_diff < 1e-2, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
mean_diff = torch.mean(torch.abs(output1 - output2))
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
# diff sum
diff_sum = torch.sum(torch.abs(output1 - output2))
logger.info(f"Diff sum between outputs: {diff_sum.item()}")
# sum
sum_output1 = torch.sum(output1.float())
sum_output2 = torch.sum(output2.float())
logger.info(f"Rank {sp_rank} Sum of output1: {sum_output1.item()}")
logger.info(f"Rank {sp_rank} Sum of output2: {sum_output2.item()}")
# Clean up
logger.info("Cleaning up distributed environment")
destroy_model_parallel()
destroy_distributed_environment()
cleanup_dist_env_and_memory()
logger.info("Test completed successfully")
if __name__ == "__main__":
test_hunyuanvideo_distributed()
-201
View File
@@ -1,201 +0,0 @@
from fastvideo.v1.models.encoders.llama import LlamaModel
from fastvideo.v1.v0_reference_src.models.hunyuan.text_encoder import TextEncoder, load_text_encoder, load_tokenizer
import os
import torch
import torch.nn as nn
import argparse
import numpy as np
from fastvideo.v1.logger import init_logger
from transformers import AutoConfig
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
logger = init_logger(__name__)
def initialize_identical_weights(model1, model2, seed=42):
"""Initialize both models with identical weights using a fixed seed for reproducibility."""
# Get all parameters from both models
params1 = dict(model1.named_parameters())
params2 = dict(model2.named_parameters())
# Initialize each layer with identical values
with torch.no_grad():
# Initialize weights
for name1, param1 in params1.items():
if 'weight' in name1:
# Set seed before each weight initialization
torch.manual_seed(seed)
nn.init.normal_(param1, mean=0.0, std=0.05)
for name2, param2 in params2.items():
if 'weight' in name2:
# Reset seed to get same initialization
torch.manual_seed(seed)
nn.init.normal_(param2, mean=0.0, std=0.05)
# Initialize biases
for name1, param1 in params1.items():
if 'bias' in name1:
torch.manual_seed(seed)
nn.init.normal_(param1, mean=0.0, std=0.05)
param1.data = param1.data.to(torch.bfloat16)
for name2, param2 in params2.items():
if 'bias' in name2:
torch.manual_seed(seed)
nn.init.normal_(param2, mean=0.0, std=0.05)
param2.data = param2.data.to(torch.bfloat16)
logger.info("Both models initialized with identical weights in bfloat16")
return model1, model2
def setup_args():
parser = argparse.ArgumentParser(description='LLaMA Encoder Test')
parser.add_argument('--model-path', type=str, default="meta-llama/Llama-2-7b-hf",
help='Path to the LLaMA model')
parser.add_argument('--precision', type=str, default="float16",
help='Precision to use for the model (float32, float16, bfloat16)')
return parser.parse_args()
def test_llama_encoder():
init_distributed_environment(world_size=1, rank=0, distributed_init_method="env://", local_rank=0, backend="nccl")
initialize_model_parallel(tensor_model_parallel_size=1, sequence_model_parallel_size=1, backend="nccl")
args = setup_args()
# Set fixed random seed for reproducibility
torch.manual_seed(42)
np.random.seed(42)
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
# Initialize the two model implementations
logger.info(f"Loading models from {args.model_path}")
model_path = "data/hunyuanvideo-community/HunyuanVideo/text_encoder"
hf_config = AutoConfig.from_pretrained(model_path)
print(hf_config)
# Load our implementation using the loader from text_encoder/__init__.py
model1, _ = load_text_encoder(
text_encoder_type="llm",
text_encoder_precision='fp16',
text_encoder_path=model_path,
logger=logger,
device=device
)
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
from fastvideo.v1.models.loader.loader import TextEncoderLoader
loader = TextEncoderLoader()
args.device_str = "cuda:0"
model2 = loader.load_model(model_path, hf_config, args)
# Convert to float16 and move to device
model2 = model2.to(torch.float16)
model2 = model2.to(device)
model2.eval()
# Sanity check weights between the two models
logger.info("Comparing model weights for sanity check...")
params1 = dict(model1.named_parameters())
params2 = dict(model2.named_parameters())
# Check number of parameters
logger.info(f"Model1 has {len(params1)} parameters")
logger.info(f"Model2 has {len(params2)} parameters")
# Compare a few key parameters
weight_diffs = []
# check if embed_tokens are the same
print(model1.embed_tokens.weight.shape, model2.embed_tokens.weight.shape)
assert torch.allclose(model1.embed_tokens.weight, model2.embed_tokens.weight)
weights = ["layers.{}.input_layernorm.weight", "layers.{}.post_attention_layernorm.weight"]
# for (name1, param1), (name2, param2) in zip(
# sorted(params1.items()), sorted(params2.items())
# ):
for l in range(hf_config.num_hidden_layers):
for w in weights:
name1 = w.format(l)
name2 = w.format(l)
p1 = params1[name1]
p2 = params2[name2]
print(type(p2))
if "gate_up" in name2:
print("skipping gate_up")
continue
try:
logger.info(f"Parameter: {name1} vs {name2}")
max_diff = torch.max(torch.abs(p1 - p2)).item()
mean_diff = torch.mean(torch.abs(p1 - p2)).item()
weight_diffs.append((name1, name2, max_diff, mean_diff))
logger.info(f" Max diff: {max_diff}, Mean diff: {mean_diff}")
except Exception as e:
logger.info(f"Error comparing {name1} and {name2}: {e}")
tokenizer_path = "data/hunyuanvideo-community/HunyuanVideo/tokenizer"
# Load tokenizer
tokenizer, _ = load_tokenizer(
tokenizer_type="llm",
tokenizer_path=tokenizer_path,
logger=logger
)
# Test with some sample prompts
prompts = [
"Once upon a time",
# "The quick brown fox jumps over",
# "In a galaxy far, far away"
]
logger.info("Testing LLaMA encoder with sample prompts")
with torch.no_grad():
for prompt in prompts:
logger.info(f"Testing prompt: '{prompt}'")
# Tokenize the prompt
tokens = tokenizer(
prompt,
padding="max_length",
max_length=128,
truncation=True,
return_tensors="pt"
).to(device)
# Get outputs from our implementation
# filter out padding input_ids
# tokens.input_ids = tokens.input_ids[tokens.attention_mask==1]
# tokens.attention_mask = tokens.attention_mask[tokens.attention_mask==1]
outputs1 = model1(
input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
output_hidden_states=True
)
print("--------------------------------")
logger.info(f"Testing model2")
# Get outputs from HuggingFace implementation
outputs2 = model2(
input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
output_hidden_states=True
)
# Compare last hidden states
last_hidden_state1 = outputs1.last_hidden_state[tokens.attention_mask==1]
last_hidden_state2 = outputs2.last_hidden_state[tokens.attention_mask==1]
assert last_hidden_state1.shape == last_hidden_state2.shape, \
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
max_diff_hidden = torch.max(torch.abs(last_hidden_state1 - last_hidden_state2))
mean_diff_hidden = torch.mean(torch.abs(last_hidden_state1 - last_hidden_state2))
logger.info(f"Maximum difference in last hidden states: {max_diff_hidden.item()}")
logger.info(f"Mean difference in last hidden states: {mean_diff_hidden.item()}")
logger.info("Test passed! Both LLaMA encoder implementations produce similar outputs.")
logger.info("Test completed successfully")
if __name__ == "__main__":
test_llama_encoder()
-163
View File
@@ -1,163 +0,0 @@
import os
import torch
import torch.nn as nn
import torch.distributed as dist
import argparse
from fastvideo.v1.logger import init_logger
from fastvideo.v1.distributed.parallel_state import (
init_distributed_environment,
initialize_model_parallel,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
destroy_model_parallel,
destroy_distributed_environment,
cleanup_dist_env_and_memory
)
from fastvideo.v1.layers.linear import (
ColumnParallelLinear,
RowParallelLinear
)
from fastvideo.v1.distributed.communication_op import (
tensor_model_parallel_all_reduce,
tensor_model_parallel_all_gather
)
logger = init_logger(__name__)
class SimpleTPModel(nn.Module):
"""A simple model that uses tensor parallelism."""
def __init__(self, hidden_size=1024, intermediate_size=4096):
super().__init__()
# Column parallel linear layer (splits output dimension)
self.fc1 = ColumnParallelLinear(
input_size=hidden_size,
output_size=intermediate_size,
bias=True,
gather_output=False, # Don't gather output since we're passing to row parallel
skip_bias_add=False
)
# Row parallel linear layer (splits input dimension)
self.fc2 = RowParallelLinear(
input_size=intermediate_size,
output_size=hidden_size,
bias=True,
input_is_parallel=True, # Input is already split from previous layer
skip_bias_add=False
)
self.activation = nn.GELU()
def forward(self, x):
# Forward through column parallel layer
hidden_states, _ = self.fc1(x)
# Apply activation
hidden_states = self.activation(hidden_states)
# Forward through row parallel layer
output, _ = self.fc2(hidden_states)
return output
def initialize_random_weights(model, seed=42):
"""Initialize the model with random weights using a fixed seed for reproducibility."""
# Set seed for reproducibility
torch.manual_seed(seed)
# Initialize weights for each layer
with torch.no_grad():
# For ColumnParallelLinear layers
if hasattr(model, 'fc1'):
nn.init.normal_(model.fc1.weight, mean=0.0, std=0.02)
if model.fc1.bias is not None:
nn.init.zeros_(model.fc1.bias)
# For RowParallelLinear layers
if hasattr(model, 'fc2'):
nn.init.normal_(model.fc2.weight, mean=0.0, std=0.02)
if model.fc2.bias is not None:
nn.init.zeros_(model.fc2.bias)
logger.info("Model initialized with random weights")
return model
def setup_args():
parser = argparse.ArgumentParser(description='Simple Tensor Parallelism Example')
parser.add_argument('--tensor-model-parallel-size', type=int, default=8,
help='Degree of tensor model parallelism')
parser.add_argument('--batch-size', type=int, default=8,
help='Batch size for the example')
parser.add_argument('--hidden-size', type=int, default=1024,
help='Hidden size for the model')
parser.add_argument('--intermediate-size', type=int, default=4096,
help='Intermediate size for the model')
return parser.parse_args()
def main():
args = setup_args()
# Initialize distributed environment
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))
logger.info(f"Initializing process: rank={rank}, local_rank={local_rank}, world_size={world_size}")
init_distributed_environment(
world_size=world_size,
rank=rank,
local_rank=local_rank
)
# Initialize tensor model parallel groups
initialize_model_parallel(
tensor_model_parallel_size=args.tensor_model_parallel_size
)
# Get tensor parallel info
tp_rank = get_tensor_model_parallel_rank()
tp_world_size = get_tensor_model_parallel_world_size()
logger.info(f"Process rank {rank} initialized with TP rank {tp_rank} in TP world size {tp_world_size}")
# Create a simple model
model = SimpleTPModel(
hidden_size=args.hidden_size,
intermediate_size=args.intermediate_size
)
# Initialize with random weights
model = initialize_random_weights(model)
# Create a random input tensor
batch_size = args.batch_size
hidden_size = args.hidden_size
x = torch.randn(batch_size, hidden_size, dtype=torch.float)
# Move to GPU if available
device = torch.device(f"cuda:{local_rank}" if torch.cuda.is_available() else "cpu")
model = model.to(device)
x = x.to(device)
# Forward pass
logger.info(f"Running forward pass on TP rank {tp_rank}")
with torch.no_grad():
output = model(x)
# Print output shape and statistics
logger.info(f"Output shape: {output.shape}")
logger.info(f"Output mean: {output.mean().item()}, std: {output.std().item()}")
# Clean up
logger.info("Cleaning up distributed environment")
destroy_model_parallel()
destroy_distributed_environment()
cleanup_dist_env_and_memory()
logger.info("Example completed successfully")
if __name__ == "__main__":
main()
-312
View File
@@ -1,312 +0,0 @@
import torch
import fastvideo.v1.envs as envs
import inspect
from fastvideo.v1.logger import init_logger
import argparse
import math
import sys
from typing import List, Dict, Union, Type, Any, TypeVar
import yaml
from functools import wraps
logger = init_logger(__name__)
T = TypeVar("T")
# TODO(will): used to convert inference_args.precision to torch.dtype. Find a
# cleaner way to do this.
PRECISION_TO_TYPE = {
"fp32": torch.float32,
"fp16": torch.float16,
"bf16": torch.bfloat16,
}
def find_nccl_library() -> str:
"""
We either use the library file specified by the `VLLM_NCCL_SO_PATH`
environment variable, or we find the library file brought by PyTorch.
After importing `torch`, `libnccl.so.2` or `librccl.so.1` can be
found by `ctypes` automatically.
"""
so_file = envs.FASTVIDEO_NCCL_SO_PATH
# manually load the nccl library
if so_file:
logger.info(
"Found nccl from environment variable FASTVIDEO_NCCL_SO_PATH=%s",
so_file)
else:
if torch.version.cuda is not None:
so_file = "libnccl.so.2"
elif torch.version.hip is not None:
so_file = "librccl.so.1"
else:
raise ValueError("NCCL only supports CUDA and ROCm backends.")
logger.info("Found nccl from library %s", so_file)
return so_file
prev_set_stream = torch.cuda.set_stream
_current_stream = None
def _patched_set_stream(stream: torch.cuda.Stream) -> None:
global _current_stream
_current_stream = stream
prev_set_stream(stream)
torch.cuda.set_stream = _patched_set_stream
def current_stream() -> torch.cuda.Stream:
"""
replace `torch.cuda.current_stream()` with `vllm.utils.current_stream()`.
it turns out that `torch.cuda.current_stream()` is quite expensive,
as it will construct a new stream object at each call.
here we patch `torch.cuda.set_stream` to keep track of the current stream
directly, so that we can avoid calling `torch.cuda.current_stream()`.
the underlying hypothesis is that we do not call `torch._C._cuda_setStream`
from C/C++ code.
"""
from fastvideo.v1.platforms import current_platform
global _current_stream
if _current_stream is None:
# when this function is called before any stream is set,
# we return the default stream.
# On ROCm using the default 0 stream in combination with RCCL
# is hurting performance. Therefore creating a dedicated stream
# per process
_current_stream = torch.cuda.Stream() if current_platform.is_rocm(
) else torch.cuda.current_stream()
return _current_stream
class StoreBoolean(argparse.Action):
def __call__(self, parser, namespace, values, option_string=None):
if values.lower() == "true":
setattr(namespace, self.dest, True)
elif values.lower() == "false":
setattr(namespace, self.dest, False)
else:
raise ValueError(f"Invalid boolean value: {values}. "
"Expected 'true' or 'false'.")
class SortedHelpFormatter(argparse.HelpFormatter):
"""SortedHelpFormatter that sorts arguments by their option strings."""
def add_arguments(self, actions):
actions = sorted(actions, key=lambda x: x.option_strings)
super().add_arguments(actions)
class FlexibleArgumentParser(argparse.ArgumentParser):
"""ArgumentParser that allows both underscore and dash in names."""
def __init__(self, *args, **kwargs):
# Set the default 'formatter_class' to SortedHelpFormatter
if 'formatter_class' not in kwargs:
kwargs['formatter_class'] = SortedHelpFormatter
super().__init__(*args, **kwargs)
def parse_args(self, args=None, namespace=None):
if args is None:
args = sys.argv[1:]
if '--config' in args:
args = self._pull_args_from_config(args)
# Convert underscores to dashes and vice versa in argument names
processed_args = []
for arg in args:
if arg.startswith('--'):
if '=' in arg:
key, value = arg.split('=', 1)
key = '--' + key[len('--'):].replace('_', '-')
processed_args.append(f'{key}={value}')
else:
processed_args.append('--' +
arg[len('--'):].replace('_', '-'))
elif arg.startswith('-O') and arg != '-O' and len(arg) == 2:
# allow -O flag to be used without space, e.g. -O3
processed_args.append('-O')
processed_args.append(arg[2:])
else:
processed_args.append(arg)
return super().parse_args(processed_args, namespace)
def _pull_args_from_config(self, args: List[str]) -> List[str]:
"""Method to pull arguments specified in the config file
into the command-line args variable.
The arguments in config file will be inserted between
the argument list.
example:
```yaml
port: 12323
tensor-parallel-size: 4
```
```python
$: vllm {serve,chat,complete} "facebook/opt-12B" \
--config config.yaml -tp 2
$: args = [
"serve,chat,complete",
"facebook/opt-12B",
'--config', 'config.yaml',
'-tp', '2'
]
$: args = [
"serve,chat,complete",
"facebook/opt-12B",
'--port', '12323',
'--tensor-parallel-size', '4',
'-tp', '2'
]
```
Please note how the config args are inserted after the sub command.
this way the order of priorities is maintained when these are args
parsed by super().
"""
assert args.count(
'--config') <= 1, "More than one config file specified!"
index = args.index('--config')
if index == len(args) - 1:
raise ValueError("No config file specified! \
Please check your command-line arguments.")
file_path = args[index + 1]
config_args = self._load_config_file(file_path)
# 0th index is for {serve,chat,complete}
# followed by model_tag (only for serve)
# followed by config args
# followed by rest of cli args.
# maintaining this order will enforce the precedence
# of cli > config > defaults
if args[0] == "serve":
if index == 1:
raise ValueError(
"No model_tag specified! Please check your command-line"
" arguments.")
args = [args[0]] + [
args[1]
] + config_args + args[2:index] + args[index + 2:]
else:
args = [args[0]] + config_args + args[1:index] + args[index + 2:]
return args
def _load_config_file(self, file_path: str) -> List[str]:
"""Loads a yaml file and returns the key value pairs as a
flattened list with argparse like pattern
```yaml
port: 12323
tensor-parallel-size: 4
```
returns:
processed_args: list[str] = [
'--port': '12323',
'--tensor-parallel-size': '4'
]
"""
extension: str = file_path.split('.')[-1]
if extension not in ('yaml', 'yml'):
raise ValueError(
"Config file must be of a yaml/yml type.\
%s supplied", extension)
# only expecting a flat dictionary of atomic types
processed_args: List[str] = []
config: Dict[str, Union[int, str]] = {}
try:
with open(file_path) as config_file:
config = yaml.safe_load(config_file)
except Exception as ex:
logger.error(
"Unable to read the config file at %s. \
Make sure path is correct", file_path)
raise ex
store_boolean_arguments = [
action.dest for action in self._actions
if isinstance(action, StoreBoolean)
]
for key, value in config.items():
if isinstance(value, bool) and key not in store_boolean_arguments:
if value:
processed_args.append('--' + key)
else:
processed_args.append('--' + key)
processed_args.append(str(value))
return processed_args
def warn_for_unimplemented_methods(cls: Type[T]) -> Type[T]:
"""
A replacement for `abc.ABC`.
When we use `abc.ABC`, subclasses will fail to instantiate
if they do not implement all abstract methods.
Here, we only require `raise NotImplementedError` in the
base class, and log a warning if the method is not implemented
in the subclass.
"""
original_init = cls.__init__
def find_unimplemented_methods(self: object):
unimplemented_methods = []
for attr_name in dir(self):
# bypass inner method
if attr_name.startswith('_'):
continue
try:
attr = getattr(self, attr_name)
# get the func of callable method
if callable(attr):
attr_func = attr.__func__
except AttributeError:
continue
src = inspect.getsource(attr_func)
if "NotImplementedError" in src:
unimplemented_methods.append(attr_name)
if unimplemented_methods:
method_names = ','.join(unimplemented_methods)
msg = (f"Methods {method_names} not implemented in {self}")
logger.warning(msg)
@wraps(original_init)
def wrapped_init(self, *args, **kwargs) -> None:
original_init(self, *args, **kwargs)
find_unimplemented_methods(self)
type.__setattr__(cls, '__init__', wrapped_init)
return cls
def align_to(value, alignment):
"""align height, width according to alignment
Args:
value (int): height or width
alignment (int): target alignment factor
Returns:
int: the aligned value
"""
return int(math.ceil(value / alignment) * alignment)
+1 -1
View File
@@ -12,7 +12,7 @@ import torch
import torchvision
from cog import BasePredictor, Input, Path
from einops import rearrange
from transformers import LlamaForCausalLM
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
MODEL_CACHE = 'FastHunyuan'
+3 -6
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "fastvideo"
version = "0.0.1.dev1"
version = "0.0.1"
description = "FastVideo"
readme = "README.md"
requires-python = ">=3.8"
@@ -39,10 +39,7 @@ dependencies = [
"gpustat", "watch",
# Kernel & Packaging
"wheel",
# Sliding Tile Atteniton Kernel
"st_attn>=0.0.1"
"wheel"
]
@@ -76,4 +73,4 @@ skip ="./data,./wandb,./csrc/sliding_tile_attention/tk"
column_limit = 120
[tool.isort]
line_length = 120
line_length = 120
-1
View File
@@ -1 +0,0 @@
huggingface_hub
+1 -1
View File
@@ -37,4 +37,4 @@ torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
--output_path outputs_video/hunyuan/vae_sp/ \
--model_path $MODEL_BASE \
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt \
--vae-sp
--vae-sp
-28
View File
@@ -1,28 +0,0 @@
#!/bin/bash
num_gpus=8
export TOKENIZERS_PARALLELISM=true
export MODEL_BASE=data/hunyuan_diffusers
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--use-v1-transformer \
--use-v1-vae \
--use-v1-text-encoder \
--sp_size $num_gpus \
--tp_size $num_gpus \
--height 768 \
--width 1280 \
--num_frames 125 \
--num_inference_steps 50 \
--guidance_scale 1 \
--embedded_cfg_scale 6 \
--flow_shift 7 \
--prompt_path ./assets/prompt.txt \
--seed 12345 \
--output_path outputs_video/ \
--model_path $MODEL_BASE \
--vae-sp