Compare commits
2
Commits
docs-build
..
quant
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
008ee2099a | ||
|
|
137f61f2fe |
@@ -19,6 +19,16 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_vae_test:
|
||||
description: "Run vae-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_transformer_test:
|
||||
description: "Run transformer-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_ssim_test:
|
||||
description: "Run ssim-test"
|
||||
required: false
|
||||
@@ -39,6 +49,8 @@ jobs:
|
||||
if: ${{ github.event.pull_request.draft == false || github.event_name == 'workflow_dispatch' }}
|
||||
outputs:
|
||||
encoder-test: ${{ steps.filter.outputs.encoder-test }}
|
||||
vae-test: ${{ steps.filter.outputs.vae-test }}
|
||||
transformer-test: ${{ steps.filter.outputs.transformer-test }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
@@ -49,6 +61,14 @@ jobs:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
vae-test:
|
||||
- 'fastvideo/v1/models/vaes/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/vaes/**'
|
||||
transformer-test:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -99,6 +119,104 @@ jobs:
|
||||
JOB_ID: "encoder-test"
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
vae-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.vae-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_vae_test == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
environment: runpod-runners
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
mkdir -p ~/.ssh
|
||||
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
|
||||
chmod 600 ~/.ssh/id_rsa
|
||||
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Run tests on RunPod
|
||||
env:
|
||||
JOB_ID: "vae-test"
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
timeout-minutes: 30
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
--gpu-type "NVIDIA A40"
|
||||
--gpu-count 1
|
||||
--volume-size 100
|
||||
--test-command "pip install -e .[test] &&
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
|
||||
pytest ./fastvideo/v1/tests/vaes -s"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
env:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
JOB_ID: "vae-test"
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
transformer-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.transformer-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_transformer_test == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
environment: runpod-runners
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
mkdir -p ~/.ssh
|
||||
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
|
||||
chmod 600 ~/.ssh/id_rsa
|
||||
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Run tests on RunPod
|
||||
env:
|
||||
JOB_ID: "transformer-test"
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
timeout-minutes: 30
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
--gpu-type "NVIDIA L40S"
|
||||
--gpu-count 1
|
||||
--volume-size 100
|
||||
--test-command "pip install -e .[test] &&
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
|
||||
pytest ./fastvideo/v1/tests/transformers -s"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
env:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
JOB_ID: "transformer-test"
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
ssim-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
@@ -130,13 +248,13 @@ jobs:
|
||||
JOB_ID: "ssim-test"
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
timeout-minutes: 30
|
||||
timeout-minutes: 45
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
--gpu-type "NVIDIA A40"
|
||||
--gpu-count 2
|
||||
--disk-size 100
|
||||
--volume-size 100
|
||||
--disk-size 200
|
||||
--volume-size 200
|
||||
--test-command "pip install -e .[test] &&
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
|
||||
pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
@@ -150,7 +268,7 @@ jobs:
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
runpod-cleanup:
|
||||
needs: [encoder-test, ssim-test] # Add other jobs to this list as you create them
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
|
||||
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
@@ -167,7 +285,7 @@ jobs:
|
||||
|
||||
- name: Cleanup all RunPod instances
|
||||
env:
|
||||
JOB_IDS: '["encoder-test", "ssim-test"]' # JSON array of job IDs
|
||||
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test"]' # JSON array of job IDs
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
@@ -40,10 +40,10 @@ repos:
|
||||
- id: codespell
|
||||
additional_dependencies: ['tomli']
|
||||
args: ['--toml', 'pyproject.toml']
|
||||
- repo: https://github.com/PyCQA/isort
|
||||
rev: 0a0b7a830386ba6a31c2ec8316849ae4d1b8240d # 6.0.0
|
||||
hooks:
|
||||
- id: isort
|
||||
# - repo: https://github.com/PyCQA/isort
|
||||
# rev: 0a0b7a830386ba6a31c2ec8316849ae4d1b8240d # 6.0.0
|
||||
# hooks:
|
||||
# - id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.29
|
||||
hooks:
|
||||
@@ -66,7 +66,7 @@ repos:
|
||||
entry: bash
|
||||
args:
|
||||
- -c
|
||||
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
language: system
|
||||
always_run: true
|
||||
pass_filenames: false
|
||||
|
||||
@@ -1,21 +0,0 @@
|
||||
# Read the Docs configuration file
|
||||
# See https://docs.readthedocs.io/en/stable/config-file/v2.html for details
|
||||
|
||||
version: 2
|
||||
|
||||
build:
|
||||
os: ubuntu-22.04
|
||||
tools:
|
||||
python: "3.12"
|
||||
|
||||
sphinx:
|
||||
configuration: docs/source/conf.py
|
||||
fail_on_warning: true
|
||||
|
||||
# If using Sphinx, optionally build your docs in additional formats such as PDF
|
||||
formats: []
|
||||
|
||||
# Optionally declare the Python requirements required to build your docs
|
||||
python:
|
||||
install:
|
||||
- requirements: docs/requirements-docs.txt
|
||||
@@ -0,0 +1,44 @@
|
||||
(wanvideo)=
|
||||
|
||||
# WanVideo
|
||||
## Inference T2V with WanVideo
|
||||
First, download the model:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-T2V-1.3B-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
|
||||
```
|
||||
|
||||
or
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-T2V-14B-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
|
||||
```
|
||||
|
||||
Then run the inference using:
|
||||
|
||||
```bash
|
||||
sh scripts/inference/v1_inference_wan.sh
|
||||
```
|
||||
|
||||
Remember to set `MODEL_BASE` and `num_gpus` accordingly.
|
||||
|
||||
## Inference I2V with WanVideo
|
||||
First, download the model:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-I2V-14B-480P-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
|
||||
```
|
||||
|
||||
or
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-I2V-14B-720P-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
|
||||
```
|
||||
|
||||
Then run the inference using:
|
||||
|
||||
```bash
|
||||
sh scripts/inference/v1_inference_wan_i2v.sh
|
||||
```
|
||||
|
||||
Remember to set `MODEL_BASE` and `num_gpus` accordingly.
|
||||
@@ -31,7 +31,7 @@ class InferenceArgs:
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: int = 7
|
||||
flow_shift: Optional[float] = None
|
||||
|
||||
output_type: str = "pil"
|
||||
|
||||
@@ -43,6 +43,9 @@ class InferenceArgs:
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_len: int = 256
|
||||
@@ -54,14 +57,14 @@ class InferenceArgs:
|
||||
|
||||
# Flow Matching parameters
|
||||
flow_solver: str = "euler"
|
||||
denoise_type: str = "flow"
|
||||
denoise_type: str = "flow" # Deprecated. Will use scheduler_config.json
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
# Scheduler options
|
||||
scheduler_type: str = "euler"
|
||||
scheduler_type: str = "euler" # Deprecated. Will use the param in scheduler_config.json
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
num_videos: int = 1
|
||||
@@ -73,6 +76,7 @@ class InferenceArgs:
|
||||
log_level: str = "info"
|
||||
|
||||
# Inference parameters
|
||||
image_path: Optional[str] = None
|
||||
prompt: Optional[str] = None
|
||||
prompt_path: Optional[str] = None
|
||||
output_path: str = "outputs/"
|
||||
@@ -187,7 +191,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--flow-shift",
|
||||
"--shift",
|
||||
type=int,
|
||||
type=float,
|
||||
default=InferenceArgs.flow_shift,
|
||||
help="Flow shift parameter",
|
||||
)
|
||||
@@ -240,6 +244,16 @@ class InferenceArgs:
|
||||
default=InferenceArgs.text_len,
|
||||
help="Maximum text length",
|
||||
)
|
||||
|
||||
# Image encoder config
|
||||
parser.add_argument(
|
||||
"--image-encoder-precision",
|
||||
type=str,
|
||||
default=InferenceArgs.image_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for image encoder",
|
||||
)
|
||||
|
||||
# Secondary text encoder
|
||||
|
||||
parser.add_argument(
|
||||
@@ -343,6 +357,10 @@ class InferenceArgs:
|
||||
help="Path to a text file containing the prompt",
|
||||
)
|
||||
|
||||
parser.add_argument("--image-path",
|
||||
type=str,
|
||||
help="Path to the image for I2V generation")
|
||||
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
type=str,
|
||||
|
||||
@@ -106,6 +106,7 @@ class InferenceEngine:
|
||||
guidance_scale = inference_args.guidance_scale
|
||||
flow_shift = inference_args.flow_shift
|
||||
embedded_guidance_scale = inference_args.embedded_cfg_scale
|
||||
image_path = inference_args.image_path
|
||||
|
||||
# ========================================================================
|
||||
# Arguments: target_width, target_height, target_video_length
|
||||
@@ -162,6 +163,7 @@ class InferenceEngine:
|
||||
# local_rank = sp_group.rank
|
||||
device = torch.device(inference_args.device_str)
|
||||
batch = ForwardBatch(
|
||||
image_path=image_path,
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
|
||||
@@ -1,15 +1,46 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
# TODO
|
||||
class BaseDiT(nn.Module):
|
||||
class BaseDiT(nn.Module, ABC):
|
||||
_fsdp_shard_conditions: list = []
|
||||
attention_head_dim: int | None = None
|
||||
_param_names_mapping: dict
|
||||
hidden_size: int
|
||||
num_attention_heads: int
|
||||
|
||||
def __init_subclass__(cls):
|
||||
required_class_attrs = [
|
||||
"_fsdp_shard_conditions", "_param_names_mapping"
|
||||
]
|
||||
super().__init_subclass__()
|
||||
for attr in required_class_attrs:
|
||||
if not hasattr(cls, attr):
|
||||
raise AttributeError(
|
||||
f"Subclasses of BaseDiT must define '{attr}' class variable"
|
||||
)
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
super().__init__()
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
@abstractmethod
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
|
||||
timestep: torch.LongTensor,
|
||||
guidance=None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
pass
|
||||
|
||||
def __post_init__(self):
|
||||
required_attrs = ["hidden_size", "num_attention_heads"]
|
||||
for attr in required_attrs:
|
||||
if not hasattr(self, attr):
|
||||
raise AttributeError(
|
||||
f"Subclasses of BaseDiT must define '{attr}' instance variable"
|
||||
)
|
||||
|
||||
@@ -648,6 +648,8 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
|
||||
self.out_channels,
|
||||
dtype=dtype)
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
# TODO: change the input the FORWAD_BACTCH Dict
|
||||
# TODO: change output to a dict
|
||||
def forward(
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import Any, Dict, Optional, Tuple, Union
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -34,9 +34,10 @@ class WanImageEmbedding(torch.nn.Module):
|
||||
|
||||
def forward(self,
|
||||
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
dtype = encoder_hidden_states_image.dtype
|
||||
hidden_states = self.norm1(encoder_hidden_states_image)
|
||||
hidden_states = self.ff(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states).to(dtype)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -52,10 +53,7 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
self.time_embedder = TimestepEmbedder(
|
||||
dim,
|
||||
frequency_embedding_size=time_freq_dim,
|
||||
act_layer="silu",
|
||||
freq_dtype=torch.float64)
|
||||
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
|
||||
self.time_modulation = ModulateProjection(dim,
|
||||
factor=6,
|
||||
act_layer="silu")
|
||||
@@ -75,9 +73,8 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
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)
|
||||
temb = self.time_embedder(timestep)
|
||||
timestep_proj = self.time_modulation(temb)
|
||||
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
if encoder_hidden_states_image is not None:
|
||||
@@ -116,7 +113,9 @@ class WanSelfAttention(nn.Module):
|
||||
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
# Scaled dot product attention
|
||||
self.attn = LocalAttention(dropout_rate=0,
|
||||
self.attn = LocalAttention(num_heads=num_heads,
|
||||
head_size=self.head_dim,
|
||||
dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
causal=False)
|
||||
|
||||
@@ -284,16 +283,15 @@ class WanTransformerBlock(nn.Module):
|
||||
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 orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb.float()
|
||||
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
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
@@ -320,6 +318,8 @@ class WanTransformerBlock(nn.Module):
|
||||
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)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -327,10 +327,13 @@ class WanTransformerBlock(nn.Module):
|
||||
context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -400,8 +403,9 @@ class WanTransformer3DModel(BaseDiT):
|
||||
super().__init__()
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
self.inner_dim = inner_dim
|
||||
self.hidden_size = inner_dim
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels or in_channels
|
||||
self.patch_size = patch_size
|
||||
self.text_len = text_len
|
||||
@@ -440,19 +444,28 @@ class WanTransformer3DModel(BaseDiT):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Union[torch.Tensor, List[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,
|
||||
encoder_hidden_states_image: Optional[Union[torch.Tensor,
|
||||
List[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)
|
||||
guidance=None,
|
||||
) -> torch.Tensor:
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
if isinstance(encoder_hidden_states_image,
|
||||
list) and len(encoder_hidden_states_image) > 0:
|
||||
encoder_hidden_states_image = encoder_hidden_states_image[0]
|
||||
else:
|
||||
encoder_hidden_states_image = None
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
@@ -461,40 +474,23 @@ class WanTransformer3DModel(BaseDiT):
|
||||
post_patch_width = width // p_w
|
||||
|
||||
# Get rotary embeddings
|
||||
d = self.inner_dim // self.num_attention_heads
|
||||
d = self.hidden_size // 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.hidden_size,
|
||||
self.num_attention_heads,
|
||||
rope_dim_list,
|
||||
dtype=torch.float64,
|
||||
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
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) 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)
|
||||
@@ -504,6 +500,7 @@ class WanTransformer3DModel(BaseDiT):
|
||||
encoder_hidden_states = torch.concat(
|
||||
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
|
||||
assert encoder_hidden_states.dtype == orig_dtype
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
@@ -516,39 +513,16 @@ class WanTransformer3DModel(BaseDiT):
|
||||
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)
|
||||
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)
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
return output.float()
|
||||
|
||||
def unpatchify(self, x, grid_sizes) -> torch.Tensor:
|
||||
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
|
||||
return output
|
||||
|
||||
@@ -598,6 +598,9 @@ class CLIPVisionTransformer(nn.Module):
|
||||
inputs_embeds=hidden_states,
|
||||
return_all_hidden_states=return_all_hidden_states)
|
||||
|
||||
if not return_all_hidden_states:
|
||||
encoder_outputs = encoder_outputs[0]
|
||||
|
||||
# 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,
|
||||
@@ -654,6 +657,8 @@ class CLIPVisionModel(nn.Module, SupportsQuant):
|
||||
layer_count = len(self.vision_model.encoder.layers)
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
if name.startswith("visual_projection"):
|
||||
continue
|
||||
# post_layernorm is not needed in CLIPVisionModel
|
||||
if (name.startswith("vision_model.post_layernorm")
|
||||
and self.vision_model.post_layernorm is None):
|
||||
|
||||
@@ -28,11 +28,12 @@ import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from transformers import T5Config
|
||||
|
||||
from fastvideo.v1.distributed import get_tensor_model_parallel_world_size
|
||||
from fastvideo.v1.distributed import (get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size)
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.layers.layernorm import RMSNorm
|
||||
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
|
||||
QKVParallelLinear, RowParallelLinear)
|
||||
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
|
||||
|
||||
@@ -67,7 +68,8 @@ class T5DenseActDense(nn.Module):
|
||||
config: T5Config,
|
||||
quant_config: Optional[QuantizationConfig] = None):
|
||||
super().__init__()
|
||||
self.wi = ColumnParallelLinear(config.d_model, config.d_ff, bias=False)
|
||||
self.wi = MergedColumnParallelLinear(config.d_model, [config.d_ff],
|
||||
bias=False)
|
||||
self.wo = RowParallelLinear(config.d_ff,
|
||||
config.d_model,
|
||||
bias=False,
|
||||
@@ -87,14 +89,12 @@ class T5DenseGatedActDense(nn.Module):
|
||||
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)
|
||||
self.wi_0 = MergedColumnParallelLinear(config.d_model, [config.d_ff],
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
self.wi_1 = MergedColumnParallelLinear(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,
|
||||
@@ -170,6 +170,7 @@ class T5Attention(nn.Module):
|
||||
config.relative_attention_max_distance
|
||||
self.d_model = config.d_model
|
||||
self.key_value_proj_dim = config.d_kv
|
||||
self.total_num_heads = self.total_num_kv_heads = config.num_heads
|
||||
|
||||
# Partition heads across multiple tensor parallel GPUs.
|
||||
tp_world_size = get_tensor_model_parallel_world_size()
|
||||
@@ -178,29 +179,33 @@ class T5Attention(nn.Module):
|
||||
|
||||
self.inner_dim = self.n_heads * self.key_value_proj_dim
|
||||
# No GQA in t5.
|
||||
self.n_kv_heads = self.n_heads
|
||||
# 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)
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
self.d_model,
|
||||
self.d_model // self.total_num_heads,
|
||||
self.total_num_heads,
|
||||
self.total_num_kv_heads,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.qkv_proj",
|
||||
)
|
||||
|
||||
self.attn = T5MultiHeadAttention()
|
||||
|
||||
if self.has_relative_attention_bias:
|
||||
self.relative_attention_bias = \
|
||||
VocabParallelEmbedding(self.relative_attention_num_buckets,
|
||||
self.n_heads,
|
||||
self.total_num_heads,
|
||||
org_num_embeddings=self.relative_attention_num_buckets,
|
||||
padding_size=self.relative_attention_num_buckets,
|
||||
quant_config=quant_config)
|
||||
self.o = RowParallelLinear(
|
||||
self.inner_dim,
|
||||
self.d_model,
|
||||
self.d_model,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.o_proj",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@@ -295,13 +300,13 @@ class T5Attention(nn.Module):
|
||||
) -> torch.Tensor:
|
||||
bs, seq_len, _ = hidden_states.shape
|
||||
num_seqs = bs
|
||||
n, c = self.n_heads, self.d_model // self.n_heads
|
||||
n, c = self.n_heads, self.d_model // self.total_num_heads
|
||||
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)
|
||||
q = q.view(bs, -1, n, c)
|
||||
k = k.view(bs, -1, n, c)
|
||||
v = v.view(bs, -1, n, c)
|
||||
q = q.reshape(bs, seq_len, n, c)
|
||||
k = k.reshape(bs, seq_len, n, c)
|
||||
v = v.reshape(bs, seq_len, n, c)
|
||||
|
||||
assert attn_metadata is not None
|
||||
attn_bias = attn_metadata.attn_bias
|
||||
@@ -325,6 +330,11 @@ class T5Attention(nn.Module):
|
||||
-1) if attention_mask.ndim == 2 else attention_mask.unsqueeze(1)
|
||||
attn_bias.masked_fill_(attention_mask == 0,
|
||||
torch.finfo(q.dtype).min)
|
||||
|
||||
if get_tensor_model_parallel_world_size() > 1:
|
||||
rank = get_tensor_model_parallel_rank()
|
||||
attn_bias = attn_bias[:, rank * self.n_heads:(rank + 1) *
|
||||
self.n_heads, :, :]
|
||||
attn_output = self.attn(q, k, v, attn_bias)
|
||||
output, _ = self.o(attn_output)
|
||||
return output
|
||||
|
||||
@@ -95,9 +95,12 @@ def get_diffusers_config(
|
||||
Returns:
|
||||
The loaded configuration.
|
||||
"""
|
||||
config_name = "config.json"
|
||||
if "scheduler" in model:
|
||||
config_name = "scheduler_config.json"
|
||||
# Check if the model path exists
|
||||
if os.path.exists(model):
|
||||
config_file = os.path.join(model, "config.json")
|
||||
config_file = os.path.join(model, config_name)
|
||||
if os.path.exists(config_file):
|
||||
try:
|
||||
# Load the config directly from the file
|
||||
|
||||
@@ -10,7 +10,7 @@ from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from transformers import AutoTokenizer, PretrainedConfig
|
||||
from transformers import AutoImageProcessor, AutoTokenizer, PretrainedConfig
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
@@ -23,6 +23,7 @@ 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.registry import ModelRegistry
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -71,6 +72,8 @@ class ComponentLoader(ABC):
|
||||
"text_encoder_2": (TextEncoderLoader, "transformers"),
|
||||
"tokenizer": (TokenizerLoader, "transformers"),
|
||||
"tokenizer_2": (TokenizerLoader, "transformers"),
|
||||
"image_processor": (ImageProcessorLoader, "transformers"),
|
||||
"image_encoder": (ImageEncoderLoader, "transformers"),
|
||||
}
|
||||
|
||||
if module_type in module_loaders:
|
||||
@@ -209,11 +212,15 @@ class TextEncoderLoader(ComponentLoader):
|
||||
|
||||
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)
|
||||
return self.load_model(model_path, model_config, target_device,
|
||||
inference_args.text_encoder_precision)
|
||||
|
||||
def load_model(self, model_path: str, model_config,
|
||||
target_device: torch.device):
|
||||
with set_default_torch_dtype(torch.float16):
|
||||
def load_model(self,
|
||||
model_path: str,
|
||||
model_config,
|
||||
target_device: torch.device,
|
||||
dtype: str = "fp16"):
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
|
||||
with target_device:
|
||||
architectures = getattr(model_config, "architectures", [])
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
|
||||
@@ -240,6 +247,40 @@ class TextEncoderLoader(ComponentLoader):
|
||||
return model.eval()
|
||||
|
||||
|
||||
class ImageEncoderLoader(TextEncoderLoader):
|
||||
|
||||
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("HF Model config: %s", 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,
|
||||
inference_args.image_encoder_precision)
|
||||
|
||||
|
||||
class ImageProcessorLoader(ComponentLoader):
|
||||
"""Loader for image processor."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
"""Load the image processor based on the model path, architecture, and inference args."""
|
||||
logger.info("Loading image processor from %s", model_path)
|
||||
|
||||
image_processor = AutoImageProcessor.from_pretrained(model_path, )
|
||||
logger.info("Loaded image processor: %s",
|
||||
image_processor.__class__.__name__)
|
||||
return image_processor
|
||||
|
||||
|
||||
class TokenizerLoader(ComponentLoader):
|
||||
"""Loader for tokenizers."""
|
||||
|
||||
@@ -265,8 +306,6 @@ class VAELoader(ComponentLoader):
|
||||
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")
|
||||
@@ -289,14 +328,6 @@ class VAELoader(ComponentLoader):
|
||||
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
|
||||
|
||||
|
||||
@@ -326,6 +357,7 @@ class TransformerLoader(ComponentLoader):
|
||||
len(safetensors_list), model_path)
|
||||
|
||||
# initialize_sequence_parallel_group(inference_args.sp_size)
|
||||
default_dtype = PRECISION_TO_TYPE[inference_args.precision]
|
||||
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s", cls_name)
|
||||
@@ -333,12 +365,16 @@ class TransformerLoader(ComponentLoader):
|
||||
init_params=model_config,
|
||||
weight_dir_list=safetensors_list,
|
||||
device=inference_args.device,
|
||||
cpu_offload=inference_args.use_cpu_offload)
|
||||
cpu_offload=inference_args.use_cpu_offload,
|
||||
default_dtype=default_dtype)
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||
|
||||
model.eval()
|
||||
dtypes = set(param.dtype for param in model.parameters())
|
||||
if len(dtypes) > 1:
|
||||
model = model.to(default_dtype)
|
||||
model = model.eval()
|
||||
return model
|
||||
|
||||
|
||||
@@ -348,21 +384,17 @@ class SchedulerLoader(ComponentLoader):
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
"""Load the scheduler based on the model path, architecture, and inference args."""
|
||||
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)
|
||||
scheduler = FlowMatchDiscreteScheduler(
|
||||
shift=inference_args.flow_shift,
|
||||
solver=inference_args.flow_solver,
|
||||
)
|
||||
logger.info("Scheduler loaded: %s", scheduler)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid denoise type: {inference_args.denoise_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")
|
||||
|
||||
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
scheduler = scheduler_cls(**config)
|
||||
if inference_args.flow_shift is not None:
|
||||
scheduler.set_shift(inference_args.flow_shift)
|
||||
|
||||
return scheduler
|
||||
|
||||
|
||||
@@ -38,6 +38,7 @@ _TEXT_ENCODER_MODELS = {
|
||||
|
||||
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
|
||||
# "HunyuanVideoTransformer3DModel": ("image_encoder", "hunyuanvideo", "HunyuanVideoImageEncoder"),
|
||||
"CLIPVisionModelWithProjection": ("encoders", "clip", "CLIPVisionModel"),
|
||||
}
|
||||
|
||||
_VAE_MODELS = {
|
||||
@@ -46,12 +47,21 @@ _VAE_MODELS = {
|
||||
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
|
||||
}
|
||||
|
||||
_SCHEDULERS = {
|
||||
"FlowMatchEulerDiscreteScheduler":
|
||||
("schedulers", "scheduling_flow_match_euler_discrete",
|
||||
"FlowMatchDiscreteScheduler"),
|
||||
"UniPCMultistepScheduler":
|
||||
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
|
||||
}
|
||||
|
||||
_FAST_VIDEO_MODELS = {
|
||||
**_TEXT_TO_VIDEO_DIT_MODELS,
|
||||
**_IMAGE_TO_VIDEO_DIT_MODELS,
|
||||
**_TEXT_ENCODER_MODELS,
|
||||
**_IMAGE_ENCODER_MODELS,
|
||||
**_VAE_MODELS,
|
||||
**_SCHEDULERS,
|
||||
}
|
||||
|
||||
_SUBPROCESS_COMMAND = [
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from diffusers.utils import BaseOutput
|
||||
|
||||
|
||||
class BaseScheduler(ABC):
|
||||
timesteps: torch.tensor
|
||||
order: int
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
# Check if subclass has defined all required properties
|
||||
required_attributes = ['timesteps', 'order']
|
||||
|
||||
for attr in required_attributes:
|
||||
if not hasattr(self, attr):
|
||||
raise AttributeError(
|
||||
f"Subclasses of BaseScheduler must define '{attr}' property"
|
||||
)
|
||||
|
||||
@abstractmethod
|
||||
def set_shift(self, shift: float) -> None:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def set_timesteps(self, *args, **kwargs) -> None:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def scale_model_input(self,
|
||||
sample: torch.Tensor,
|
||||
timestep: Optional[int] = None) -> torch.Tensor:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
return_dict: bool = True,
|
||||
**kwargs,
|
||||
) -> Union[BaseOutput, Tuple]:
|
||||
pass
|
||||
@@ -27,6 +27,8 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
|
||||
from fastvideo.v1.models.schedulers.base import BaseScheduler
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
@@ -44,7 +46,7 @@ class FlowMatchDiscreteSchedulerOutput(BaseOutput):
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
|
||||
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
"""
|
||||
Euler scheduler.
|
||||
|
||||
@@ -74,6 +76,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
reverse: bool = True,
|
||||
solver: str = "euler",
|
||||
n_tokens: Optional[int] = None,
|
||||
**kwargs,
|
||||
):
|
||||
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
|
||||
|
||||
@@ -94,6 +97,8 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
)
|
||||
|
||||
BaseScheduler.__init__(self)
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
"""
|
||||
@@ -170,6 +175,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
|
||||
return idx
|
||||
|
||||
def set_shift(self, shift: float) -> None:
|
||||
self.config.shift = shift
|
||||
|
||||
def _init_step_index(self, timestep) -> None:
|
||||
if self.begin_index is None:
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -2,11 +2,12 @@
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from math import prod
|
||||
from typing import Iterator, Optional, Tuple
|
||||
from typing import Iterator, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size)
|
||||
@@ -20,8 +21,11 @@ class ParallelTiledVAE(ABC):
|
||||
tile_sample_stride_width: int
|
||||
tile_sample_stride_num_frames: int
|
||||
use_tiling: bool
|
||||
use_temporal_tiling: bool
|
||||
use_parallel_tiling: bool
|
||||
temporal_compression_ratio: int
|
||||
spatial_compression_ratio: int
|
||||
scaling_factor: Union[float, torch.tensor]
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
# Check if subclass has defined all required properties
|
||||
@@ -30,7 +34,8 @@ class ParallelTiledVAE(ABC):
|
||||
'tile_sample_min_num_frames', 'tile_sample_stride_height',
|
||||
'tile_sample_stride_width', 'tile_sample_stride_num_frames',
|
||||
'spatial_compression_ratio', 'temporal_compression_ratio',
|
||||
'use_tiling'
|
||||
'use_tiling', 'use_temporal_tiling', 'use_parallel_tiling',
|
||||
'scaling_factor'
|
||||
]
|
||||
|
||||
for attr in required_attributes:
|
||||
@@ -52,13 +57,13 @@ class ParallelTiledVAE(ABC):
|
||||
latent_num_frames = (num_frames -
|
||||
1) // self.temporal_compression_ratio + 1
|
||||
|
||||
if self.use_tiling and num_frames > self.tile_sample_min_num_frames:
|
||||
if self.use_tiling and self.use_temporal_tiling and num_frames > self.tile_sample_min_num_frames:
|
||||
latents = self.tiled_encode(x)[:, :, :latent_num_frames]
|
||||
elif self.use_tiling and (width > self.tile_sample_min_width
|
||||
or height > self.tile_sample_min_height):
|
||||
latents = self.spatial_tiled_encode(x)
|
||||
latents = self.spatial_tiled_encode(x)[:, :, :latent_num_frames]
|
||||
else:
|
||||
latents = self._encode(x)
|
||||
latents = self._encode(x)[:, :, :latent_num_frames]
|
||||
return DiagonalGaussianDistribution(latents)
|
||||
|
||||
def decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
@@ -69,16 +74,17 @@ class ParallelTiledVAE(ABC):
|
||||
num_sample_frames = (num_frames -
|
||||
1) * self.temporal_compression_ratio + 1
|
||||
|
||||
if self.use_tiling and get_sequence_model_parallel_world_size() > 1:
|
||||
if self.use_tiling and self.use_parallel_tiling and get_sequence_model_parallel_world_size(
|
||||
) > 1:
|
||||
return self.parallel_tiled_decode(z)[:, :, :num_sample_frames]
|
||||
if self.use_tiling and num_frames > tile_latent_min_num_frames:
|
||||
if self.use_tiling and self.use_temporal_tiling and num_frames > tile_latent_min_num_frames:
|
||||
return self.tiled_decode(z)[:, :, :num_sample_frames]
|
||||
|
||||
if self.use_tiling and (width > tile_latent_min_width
|
||||
or height > tile_latent_min_height):
|
||||
return self.spatial_tiled_decode(z)
|
||||
return self.spatial_tiled_decode(z)[:, :, :num_sample_frames]
|
||||
|
||||
return self._decode(z)
|
||||
return self._decode(z)[:, :, :num_sample_frames]
|
||||
|
||||
def blend_v(self, a: torch.Tensor, b: torch.Tensor,
|
||||
blend_extent: int) -> torch.Tensor:
|
||||
@@ -462,7 +468,7 @@ class DiagonalGaussianDistribution:
|
||||
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(
|
||||
sample = randn_tensor(
|
||||
self.mean.shape,
|
||||
generator=generator,
|
||||
device=self.parameters.device,
|
||||
|
||||
@@ -847,6 +847,9 @@ class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
|
||||
# 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
|
||||
self.use_temporal_tiling = True
|
||||
self.use_parallel_tiling = True
|
||||
self.scaling_factor = scaling_factor
|
||||
|
||||
# The minimal tile height and width for spatial tiling to be used
|
||||
self.tile_sample_min_height = 256
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
@@ -20,8 +22,8 @@ import torch.nn.functional as F
|
||||
import torch.utils.checkpoint
|
||||
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.models.utils import auto_attributes
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.utils import auto_attributes
|
||||
|
||||
CACHE_T = 2
|
||||
|
||||
@@ -647,6 +649,10 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
self.temperal_upsample = list(temperal_downsample)[::-1]
|
||||
self.latents_mean = list(latents_mean)
|
||||
self.latents_std = list(latents_std)
|
||||
self.scaling_factor = 1.0 / torch.tensor(self.config.latents_std).view(
|
||||
1, self.config.z_dim, 1, 1, 1)
|
||||
self.shift_factor = torch.tensor(self.config.latents_mean).view(
|
||||
1, self.config.z_dim, 1, 1, 1)
|
||||
|
||||
if load_encoder:
|
||||
self.encoder = WanEncoder3d(base_dim, z_dim * 2, dim_mult,
|
||||
@@ -661,6 +667,8 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
self.temperal_upsample, dropout)
|
||||
|
||||
self.use_tiling = True
|
||||
self.use_temporal_tiling = False
|
||||
self.use_parallel_tiling = False
|
||||
self.spatial_compression_ratio = 8
|
||||
self.temporal_compression_ratio = 4
|
||||
|
||||
@@ -691,12 +699,16 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
enc = torch.cat([first_frame, enc], dim=2)
|
||||
return enc
|
||||
|
||||
def spatial_tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
first_frame = x[:, :, 0, :, :].unsqueeze(2)
|
||||
first_frame = self._encode(first_frame, first_frame=True)
|
||||
|
||||
enc = ParallelTiledVAE.spatial_tiled_encode(self, x)
|
||||
enc = enc[:, :, 1:]
|
||||
enc = torch.cat([first_frame, enc], dim=2)
|
||||
return enc
|
||||
|
||||
def _decode(self, z: torch.Tensor, first_frame=False) -> torch.Tensor:
|
||||
latents_mean = (torch.tensor(self.latents_mean).view(
|
||||
1, self.z_dim, 1, 1, 1).to(z.device, z.dtype))
|
||||
latents_std = 1.0 / torch.tensor(self.latents_std).view(
|
||||
1, self.z_dim, 1, 1, 1).to(z.device, z.dtype)
|
||||
z = z / latents_std + latents_mean
|
||||
x = self.post_quant_conv(z)
|
||||
out = self.decoder(x, first_frame=first_frame)
|
||||
|
||||
@@ -711,6 +723,12 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
dec = dec[:, :, start_frame_idx:]
|
||||
return dec
|
||||
|
||||
def spatial_tiled_decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
dec = ParallelTiledVAE.spatial_tiled_decode(self, z)
|
||||
start_frame_idx = self.temporal_compression_ratio - 1
|
||||
dec = dec[:, :, start_frame_idx:]
|
||||
return dec
|
||||
|
||||
def parallel_tiled_decode(self, z: torch.FloatTensor) -> torch.FloatTensor:
|
||||
self.blend_num_frames *= 2
|
||||
dec = ParallelTiledVAE.parallel_tiled_decode(self, z)
|
||||
|
||||
@@ -0,0 +1,220 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
from typing import Callable, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import PIL.ImageOps
|
||||
import requests
|
||||
import torch
|
||||
from packaging import version
|
||||
|
||||
if version.parse(version.parse(
|
||||
PIL.__version__).base_version) >= version.parse("9.1.0"):
|
||||
PIL_INTERPOLATION = {
|
||||
"linear": PIL.Image.Resampling.BILINEAR,
|
||||
"bilinear": PIL.Image.Resampling.BILINEAR,
|
||||
"bicubic": PIL.Image.Resampling.BICUBIC,
|
||||
"lanczos": PIL.Image.Resampling.LANCZOS,
|
||||
"nearest": PIL.Image.Resampling.NEAREST,
|
||||
}
|
||||
else:
|
||||
PIL_INTERPOLATION = {
|
||||
"linear": PIL.Image.LINEAR,
|
||||
"bilinear": PIL.Image.BILINEAR,
|
||||
"bicubic": PIL.Image.BICUBIC,
|
||||
"lanczos": PIL.Image.LANCZOS,
|
||||
"nearest": PIL.Image.NEAREST,
|
||||
}
|
||||
|
||||
|
||||
def pil_to_numpy(
|
||||
images: Union[List[PIL.Image.Image], PIL.Image.Image]) -> np.ndarray:
|
||||
r"""
|
||||
Convert a PIL image or a list of PIL images to NumPy arrays.
|
||||
|
||||
Args:
|
||||
images (`PIL.Image.Image` or `List[PIL.Image.Image]`):
|
||||
The PIL image or list of images to convert to NumPy format.
|
||||
|
||||
Returns:
|
||||
`np.ndarray`:
|
||||
A NumPy array representation of the images.
|
||||
"""
|
||||
if not isinstance(images, list):
|
||||
images = [images]
|
||||
images = [np.array(image).astype(np.float32) / 255.0 for image in images]
|
||||
images_arr: np.ndarray = np.stack(images, axis=0)
|
||||
|
||||
return images_arr
|
||||
|
||||
|
||||
def numpy_to_pt(images: np.ndarray) -> torch.Tensor:
|
||||
r"""
|
||||
Convert a NumPy image to a PyTorch tensor.
|
||||
|
||||
Args:
|
||||
images (`np.ndarray`):
|
||||
The NumPy image array to convert to PyTorch format.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
A PyTorch tensor representation of the images.
|
||||
"""
|
||||
if images.ndim == 3:
|
||||
images = images[..., None]
|
||||
|
||||
images = torch.from_numpy(images.transpose(0, 3, 1, 2))
|
||||
return images
|
||||
|
||||
|
||||
def normalize(
|
||||
images: Union[np.ndarray,
|
||||
torch.Tensor]) -> Union[np.ndarray, torch.Tensor]:
|
||||
r"""
|
||||
Normalize an image array to [-1,1].
|
||||
|
||||
Args:
|
||||
images (`np.ndarray` or `torch.Tensor`):
|
||||
The image array to normalize.
|
||||
|
||||
Returns:
|
||||
`np.ndarray` or `torch.Tensor`:
|
||||
The normalized image array.
|
||||
"""
|
||||
return 2.0 * images - 1.0
|
||||
|
||||
|
||||
def load_image(
|
||||
image: Union[str, PIL.Image.Image],
|
||||
convert_method: Optional[Callable[[PIL.Image.Image],
|
||||
PIL.Image.Image]] = None
|
||||
) -> PIL.Image.Image:
|
||||
"""
|
||||
Loads `image` to a PIL Image.
|
||||
|
||||
Args:
|
||||
image (`str` or `PIL.Image.Image`):
|
||||
The image to convert to the PIL Image format.
|
||||
convert_method (Callable[[PIL.Image.Image], PIL.Image.Image], *optional*):
|
||||
A conversion method to apply to the image after loading it. When set to `None` the image will be converted
|
||||
"RGB".
|
||||
|
||||
Returns:
|
||||
`PIL.Image.Image`:
|
||||
A PIL Image.
|
||||
"""
|
||||
if isinstance(image, str):
|
||||
if image.startswith("http://") or image.startswith("https://"):
|
||||
image = PIL.Image.open(requests.get(image, stream=True).raw)
|
||||
elif os.path.isfile(image):
|
||||
image = PIL.Image.open(image)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Incorrect path or URL. URLs must start with `http://` or `https://`, and {image} is not a valid path."
|
||||
)
|
||||
elif isinstance(image, PIL.Image.Image):
|
||||
image = image
|
||||
else:
|
||||
raise ValueError(
|
||||
"Incorrect format used for the image. Should be a URL linking to an image, a local path, or a PIL image."
|
||||
)
|
||||
|
||||
image = PIL.ImageOps.exif_transpose(image)
|
||||
|
||||
if convert_method is not None:
|
||||
image = convert_method(image)
|
||||
else:
|
||||
image = image.convert("RGB")
|
||||
|
||||
return image
|
||||
|
||||
|
||||
def get_default_height_width(
|
||||
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
|
||||
vae_scale_factor: int,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
) -> Tuple[int, int]:
|
||||
r"""
|
||||
Returns the height and width of the image, downscaled to the next integer multiple of `vae_scale_factor`.
|
||||
|
||||
Args:
|
||||
image (`Union[PIL.Image.Image, np.ndarray, torch.Tensor]`):
|
||||
The image input, which can be a PIL image, NumPy array, or PyTorch tensor. If it is a NumPy array, it
|
||||
should have shape `[batch, height, width]` or `[batch, height, width, channels]`. If it is a PyTorch
|
||||
tensor, it should have shape `[batch, channels, height, width]`.
|
||||
height (`Optional[int]`, *optional*, defaults to `None`):
|
||||
The height of the preprocessed image. If `None`, the height of the `image` input will be used.
|
||||
width (`Optional[int]`, *optional*, defaults to `None`):
|
||||
The width of the preprocessed image. If `None`, the width of the `image` input will be used.
|
||||
|
||||
Returns:
|
||||
`Tuple[int, int]`:
|
||||
A tuple containing the height and width, both resized to the nearest integer multiple of
|
||||
`vae_scale_factor`.
|
||||
"""
|
||||
|
||||
if height is None:
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
height = image.height
|
||||
elif isinstance(image, torch.Tensor):
|
||||
height = image.shape[2]
|
||||
else:
|
||||
height = image.shape[1]
|
||||
|
||||
if width is None:
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
width = image.width
|
||||
elif isinstance(image, torch.Tensor):
|
||||
width = image.shape[3]
|
||||
else:
|
||||
width = image.shape[2]
|
||||
|
||||
width, height = (x - x % vae_scale_factor for x in (width, height)
|
||||
) # resize to integer multiple of vae_scale_factor
|
||||
|
||||
return height, width
|
||||
|
||||
|
||||
def resize(
|
||||
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
|
||||
height: int,
|
||||
width: int,
|
||||
resize_mode: str = "default", # "default", "fill", "crop"
|
||||
resample: str = "lanczos",
|
||||
) -> Union[PIL.Image.Image, np.ndarray, torch.Tensor]:
|
||||
"""
|
||||
Resize image.
|
||||
|
||||
Args:
|
||||
image (`PIL.Image.Image`, `np.ndarray` or `torch.Tensor`):
|
||||
The image input, can be a PIL image, numpy array or pytorch tensor.
|
||||
height (`int`):
|
||||
The height to resize to.
|
||||
width (`int`):
|
||||
The width to resize to.
|
||||
resize_mode (`str`, *optional*, defaults to `default`):
|
||||
The resize mode to use, can be one of `default` or `fill`. If `default`, will resize the image to fit
|
||||
within the specified width and height, and it may not maintaining the original aspect ratio. If `fill`,
|
||||
will resize the image to fit within the specified width and height, maintaining the aspect ratio, and
|
||||
then center the image within the dimensions, filling empty with data from image. If `crop`, will resize
|
||||
the image to fit within the specified width and height, maintaining the aspect ratio, and then center
|
||||
the image within the dimensions, cropping the excess. Note that resize_mode `fill` and `crop` are only
|
||||
supported for PIL image input.
|
||||
|
||||
Returns:
|
||||
`PIL.Image.Image`, `np.ndarray` or `torch.Tensor`:
|
||||
The resized image.
|
||||
"""
|
||||
if resize_mode != "default" and not isinstance(image, PIL.Image.Image):
|
||||
raise ValueError(
|
||||
f"Only PIL image input is supported for resize_mode {resize_mode}")
|
||||
assert isinstance(image, PIL.Image.Image)
|
||||
if resize_mode == "default":
|
||||
image = image.resize((width, height),
|
||||
resample=PIL_INTERPOLATION[resample])
|
||||
else:
|
||||
raise ValueError(f"resize_mode {resize_mode} is not supported")
|
||||
return image
|
||||
@@ -29,6 +29,10 @@ class ForwardBatch:
|
||||
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None
|
||||
|
||||
# Image inputs
|
||||
image_path: Optional[str] = None
|
||||
image_embeds: List[torch.Tensor] = field(default_factory=list)
|
||||
|
||||
# Text inputs
|
||||
prompt: Optional[Union[str, List[str]]] = None
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None
|
||||
@@ -55,6 +59,7 @@ class ForwardBatch:
|
||||
# Latent tensors
|
||||
latents: Optional[torch.Tensor] = None
|
||||
noise_pred: Optional[torch.Tensor] = None
|
||||
image_latent: Optional[torch.Tensor] = None
|
||||
|
||||
# Latent dimensions
|
||||
num_channels_latents: Optional[int] = None
|
||||
@@ -100,3 +105,4 @@ class ForwardBatch:
|
||||
# Set do_classifier_free_guidance based on guidance scale and negative prompt
|
||||
if self.guidance_scale > 1.0:
|
||||
self.do_classifier_free_guidance = True
|
||||
self.negative_prompt_embeds = []
|
||||
|
||||
@@ -7,15 +7,19 @@ complete diffusion pipelines.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.clip_image_encoding import (
|
||||
CLIPImageEncodingStage)
|
||||
from fastvideo.v1.pipelines.stages.clip_text_encoding import (
|
||||
CLIPTextEncodingStage)
|
||||
from fastvideo.v1.pipelines.stages.conditioning import ConditioningStage
|
||||
from fastvideo.v1.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.v1.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.v1.pipelines.stages.encoding import EncodingStage
|
||||
from fastvideo.v1.pipelines.stages.input_validation import InputValidationStage
|
||||
from fastvideo.v1.pipelines.stages.latent_preparation import (
|
||||
LatentPreparationStage)
|
||||
from fastvideo.v1.pipelines.stages.llama_encoding import LlamaEncodingStage
|
||||
from fastvideo.v1.pipelines.stages.t5_encoding import T5EncodingStage
|
||||
from fastvideo.v1.pipelines.stages.timestep_preparation import (
|
||||
TimestepPreparationStage)
|
||||
|
||||
@@ -26,7 +30,10 @@ __all__ = [
|
||||
"LatentPreparationStage",
|
||||
"ConditioningStage",
|
||||
"DenoisingStage",
|
||||
"EncodingStage",
|
||||
"DecodingStage",
|
||||
"LlamaEncodingStage",
|
||||
"T5EncodingStage",
|
||||
"CLIPTextEncodingStage",
|
||||
"CLIPImageEncodingStage",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Image encoding stages for I2V diffusion pipelines.
|
||||
|
||||
This module contains implementations of image encoding stages for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import load_image
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CLIPImageEncodingStage(PipelineStage):
|
||||
"""
|
||||
Stage for encoding image prompts into embeddings for diffusion models.
|
||||
|
||||
This stage handles the encoding of image prompts into the embedding space
|
||||
expected by the diffusion model.
|
||||
"""
|
||||
|
||||
def __init__(self, image_encoder, image_processor) -> None:
|
||||
"""
|
||||
Initialize the prompt encoding stage.
|
||||
|
||||
Args:
|
||||
enable_logging: Whether to enable logging for this stage.
|
||||
is_secondary: Whether this is a secondary image encoder.
|
||||
"""
|
||||
super().__init__()
|
||||
self.image_processor = image_processor
|
||||
self.image_encoder = image_encoder
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into image encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if inference_args.use_cpu_offload:
|
||||
self.image_encoder = self.image_encoder.to(batch.device)
|
||||
|
||||
image = load_image(batch.image_path)
|
||||
|
||||
image_inputs = self.image_processor(
|
||||
images=image, return_tensors="pt").to(batch.device)
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
image_embeds = self.image_encoder(**image_inputs)
|
||||
|
||||
batch.image_embeds.append(image_embeds)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
self.image_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
@@ -5,6 +5,8 @@ Prompt encoding stages for diffusion pipelines.
|
||||
This module contains implementations of prompt encoding stages for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -49,6 +51,8 @@ class CLIPTextEncodingStage(PipelineStage):
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if inference_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
batch.prompt,
|
||||
@@ -64,4 +68,24 @@ class CLIPTextEncodingStage(PipelineStage):
|
||||
|
||||
batch.prompt_embeds.append(prompt_embeds)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
negative_text_inputs = self.tokenizer(
|
||||
batch.negative_prompt,
|
||||
truncation=True,
|
||||
# better way to handle this?
|
||||
max_length=77,
|
||||
return_tensors="pt",
|
||||
)
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
negative_outputs = self.text_encoder(
|
||||
input_ids=negative_text_inputs["input_ids"].to(
|
||||
batch.device), )
|
||||
negative_prompt_embeds = negative_outputs["pooler_output"]
|
||||
|
||||
batch.negative_prompt_embeds.append(negative_prompt_embeds)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
@@ -39,8 +39,7 @@ class ConditioningStage(PipelineStage):
|
||||
if not batch.do_classifier_free_guidance:
|
||||
return batch
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"Classifier-free guidance is not supported yet")
|
||||
return batch
|
||||
|
||||
logger.info("batch.negative_prompt_embeds: %s",
|
||||
batch.negative_prompt_embeds)
|
||||
|
||||
@@ -54,13 +54,20 @@ class DecodingStage(PipelineStage):
|
||||
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)
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.config.scaling_factor
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
# Decode latents
|
||||
with torch.autocast(device_type="cuda",
|
||||
@@ -70,6 +77,8 @@ class DecodingStage(PipelineStage):
|
||||
self.vae.enable_tiling()
|
||||
# if inference_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
|
||||
@@ -53,6 +53,9 @@ class DenoisingStage(PipelineStage):
|
||||
Returns:
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
# If use cpu offload, need to load the model back into gpu again
|
||||
if inference_args.use_cpu_offload:
|
||||
self.transformer = self.transformer.to(batch.device)
|
||||
# Prepare extra step kwargs for scheduler
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.step,
|
||||
@@ -77,6 +80,12 @@ class DenoisingStage(PipelineStage):
|
||||
n=world_size).contiguous()
|
||||
latents = latents[:, :, rank, :, :, :]
|
||||
batch.latents = latents
|
||||
if batch.image_latent is not None:
|
||||
image_latent = rearrange(batch.image_latent,
|
||||
"b t (n s) h w -> b t n s h w",
|
||||
n=world_size).contiguous()
|
||||
image_latent = image_latent[:, :, rank, :, :, :]
|
||||
batch.image_latent = image_latent
|
||||
|
||||
# Get timesteps and calculate warmup steps
|
||||
timesteps = batch.timesteps
|
||||
@@ -98,9 +107,28 @@ class DenoisingStage(PipelineStage):
|
||||
result[t][layer][h] = value
|
||||
return result
|
||||
|
||||
# Prepare image latents and embeddings for I2V generation
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
assert torch.isnan(image_embeds[0]).sum() == 0
|
||||
image_embeds = [
|
||||
image_embed.to(target_dtype) for image_embed in image_embeds
|
||||
]
|
||||
|
||||
image_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_image": image_embeds,
|
||||
},
|
||||
)
|
||||
|
||||
# Get latents and embeddings
|
||||
latents = batch.latents
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
if batch.do_classifier_free_guidance:
|
||||
neg_prompt_embeds = batch.negative_prompt_embeds
|
||||
assert torch.isnan(neg_prompt_embeds[0]).sum() == 0
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
@@ -109,10 +137,13 @@ class DenoisingStage(PipelineStage):
|
||||
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)
|
||||
# Expand latents for I2V
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if batch.image_latent is not None:
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.image_latent],
|
||||
dim=1).to(target_dtype)
|
||||
assert torch.isnan(latent_model_input).sum() == 0
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
|
||||
@@ -160,8 +191,6 @@ class DenoisingStage(PipelineStage):
|
||||
inference_args=inference_args,
|
||||
)
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
attn_metadata = None
|
||||
else:
|
||||
attn_metadata = None
|
||||
|
||||
@@ -180,29 +209,43 @@ class DenoisingStage(PipelineStage):
|
||||
prompt_embeds,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
)
|
||||
|
||||
# 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
|
||||
if batch.do_classifier_free_guidance:
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
# inference_args=inference_args
|
||||
):
|
||||
# Run transformer
|
||||
noise_pred_uncond = self.transformer(
|
||||
latent_model_input,
|
||||
neg_prompt_embeds,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
)
|
||||
noise_pred_text = noise_pred
|
||||
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,
|
||||
)
|
||||
# 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]
|
||||
# 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 (
|
||||
@@ -218,11 +261,15 @@ class DenoisingStage(PipelineStage):
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
self.transformer.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
def prepare_extra_func_kwargs(self, func, kwargs) -> Dict[str, Any]:
|
||||
"""
|
||||
Prepare extra kwargs for the scheduler step.
|
||||
Prepare extra kwargs for the scheduler step / denoise step.
|
||||
|
||||
Args:
|
||||
func: The function to prepare kwargs for.
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Encoding stage for diffusion pipelines.
|
||||
"""
|
||||
from typing import Optional
|
||||
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
load_image, normalize,
|
||||
numpy_to_pt, pil_to_numpy, resize)
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class EncodingStage(PipelineStage):
|
||||
"""
|
||||
Stage for encoding pixel representations into latent space.
|
||||
|
||||
This stage handles the encoding of pixel representations into the final
|
||||
input format (e.g., latents).
|
||||
"""
|
||||
|
||||
def __init__(self, vae) -> None:
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode pixel representations into latent space.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded outputs.
|
||||
"""
|
||||
image_path = batch.image_path
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
if image_path is None:
|
||||
raise ValueError("Image Path must be provided")
|
||||
latent_height = batch.height // self.vae.spatial_compression_ratio
|
||||
latent_width = batch.width // self.vae.spatial_compression_ratio
|
||||
|
||||
image = load_image(image_path)
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=batch.height,
|
||||
width=batch.width).to(batch.device, dtype=torch.float32)
|
||||
image = image.unsqueeze(2)
|
||||
video_condition = torch.cat([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1],
|
||||
inference_args.num_frames - 1, batch.height,
|
||||
batch.width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=batch.device,
|
||||
dtype=torch.float32)
|
||||
|
||||
# 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
|
||||
|
||||
# Encode Image
|
||||
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()
|
||||
if not vae_autocast_enabled:
|
||||
video_condition = video_condition.to(vae_dtype)
|
||||
encoder_output = self.vae.encode(video_condition)
|
||||
|
||||
generator = batch.generator
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator[0])
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latent_condition -= self.vae.shift_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition -= self.vae.shift_factor
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latent_condition = latent_condition * self.vae.scaling_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition = latent_condition * self.vae.scaling_factor
|
||||
|
||||
mask_lat_size = torch.ones(1, 1, inference_args.num_frames,
|
||||
latent_height, latent_width)
|
||||
mask_lat_size[:, :, list(range(1, inference_args.num_frames))] = 0
|
||||
first_frame_mask = mask_lat_size[:, :, 0:1]
|
||||
first_frame_mask = torch.repeat_interleave(
|
||||
first_frame_mask,
|
||||
dim=2,
|
||||
repeats=self.vae.temporal_compression_ratio)
|
||||
mask_lat_size = torch.concat(
|
||||
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
|
||||
mask_lat_size = mask_lat_size.view(1, -1,
|
||||
self.vae.temporal_compression_ratio,
|
||||
latent_height, latent_width)
|
||||
mask_lat_size = mask_lat_size.transpose(1, 2)
|
||||
mask_lat_size = mask_lat_size.to(latent_condition.device)
|
||||
|
||||
batch.image_latent = torch.concat([mask_lat_size, latent_condition],
|
||||
dim=1)
|
||||
|
||||
# Offload models if needed
|
||||
if hasattr(self, 'maybe_free_model_hooks'):
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
return batch
|
||||
|
||||
def retrieve_latents(self,
|
||||
encoder_output: torch.Tensor,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
sample_mode: str = "sample"):
|
||||
if sample_mode == "sample":
|
||||
return encoder_output.sample(generator)
|
||||
elif sample_mode == "argmax":
|
||||
return encoder_output.mode()
|
||||
else:
|
||||
raise AttributeError(
|
||||
"Could not access latents of provided encoder_output")
|
||||
|
||||
def preprocess(
|
||||
self,
|
||||
image: PIL.Image.Image,
|
||||
vae_scale_factor: int,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
resize_mode: str = "default", # "default", "fill", "crop"
|
||||
) -> torch.Tensor:
|
||||
image = [image]
|
||||
|
||||
height, width = get_default_height_width(image[0], vae_scale_factor,
|
||||
height, width)
|
||||
image = [
|
||||
resize(i, height, width, resize_mode=resize_mode) for i in image
|
||||
]
|
||||
image = pil_to_numpy(image) # to np
|
||||
image = numpy_to_pt(image) # to pt
|
||||
|
||||
do_normalize = True
|
||||
if image.min() < 0:
|
||||
do_normalize = False
|
||||
if do_normalize:
|
||||
image = normalize(image)
|
||||
|
||||
return image
|
||||
@@ -6,6 +6,7 @@ from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
|
||||
@@ -20,9 +21,10 @@ class LatentPreparationStage(PipelineStage):
|
||||
denoised during the diffusion process.
|
||||
"""
|
||||
|
||||
def __init__(self, scheduler) -> None:
|
||||
def __init__(self, scheduler, vae=None) -> None:
|
||||
super().__init__()
|
||||
self.scheduler = scheduler
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -42,7 +44,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
# Adjust video length based on VAE version if needed
|
||||
if hasattr(self, 'adjust_video_length'):
|
||||
batch = self.adjust_video_length(batch, inference_args)
|
||||
batch = self.adjust_video_length(self.vae, batch, inference_args)
|
||||
# Determine batch size
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
@@ -101,7 +103,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
return batch
|
||||
|
||||
def adjust_video_length(self, batch: ForwardBatch,
|
||||
def adjust_video_length(self, vae: ParallelTiledVAE, batch: ForwardBatch,
|
||||
inference_args: InferenceArgs) -> ForwardBatch:
|
||||
"""
|
||||
Adjust video length based on VAE version.
|
||||
@@ -114,6 +116,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
The batch with adjusted video length.
|
||||
"""
|
||||
video_length = batch.num_frames
|
||||
temporal_scale_factor = vae.temporal_compression_ratio if vae is not None else 4
|
||||
# TODO
|
||||
batch.num_frames = (video_length - 1) // 4 + 1
|
||||
batch.num_frames = (video_length - 1) // temporal_scale_factor + 1
|
||||
return batch
|
||||
|
||||
@@ -7,6 +7,8 @@ This module contains implementations of prompt encoding stages for diffusion pip
|
||||
|
||||
from typing import TypedDict
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -71,6 +73,8 @@ class LlamaEncodingStage(PipelineStage):
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if inference_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text = prompt_template_video["template"].format(batch.prompt)
|
||||
text_inputs = self.tokenizer(
|
||||
@@ -93,4 +97,33 @@ class LlamaEncodingStage(PipelineStage):
|
||||
last_hidden_state = last_hidden_state[:, crop_start:]
|
||||
batch.prompt_embeds.append(last_hidden_state)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
negative_text = prompt_template_video["template"].format(
|
||||
batch.negative_prompt)
|
||||
negative_text_inputs = self.tokenizer(
|
||||
negative_text,
|
||||
truncation=True,
|
||||
# better way to handle this?
|
||||
max_length=256,
|
||||
return_tensors="pt",
|
||||
)
|
||||
hidden_state_skip_layer = 2
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
negative_outputs = self.text_encoder(
|
||||
input_ids=negative_text_inputs["input_ids"].to(
|
||||
batch.device),
|
||||
output_hidden_states=hidden_state_skip_layer is not None,
|
||||
)
|
||||
|
||||
negative_last_hidden_state = negative_outputs.hidden_states[-(
|
||||
hidden_state_skip_layer + 1)]
|
||||
crop_start = prompt_template_video.get("crop_start", -1)
|
||||
negative_last_hidden_state = negative_last_hidden_state[:,
|
||||
crop_start:]
|
||||
batch.negative_prompt_embeds.append(negative_last_hidden_state)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Prompt encoding stages for diffusion pipelines.
|
||||
|
||||
This module contains implementations of prompt encoding stages for diffusion pipelines.
|
||||
"""
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class T5EncodingStage(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, text_encoder, tokenizer) -> None:
|
||||
"""
|
||||
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__()
|
||||
self.text_encoder = text_encoder
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def forward(
|
||||
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 inference_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text = batch.prompt
|
||||
text_inputs = self.tokenizer(
|
||||
text,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
max_length=512,
|
||||
add_special_tokens=True,
|
||||
return_attention_mask=True,
|
||||
return_tensors="pt",
|
||||
).to(batch.device)
|
||||
text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = self.text_encoder(
|
||||
input_ids=text_input_ids,
|
||||
attention_mask=mask,
|
||||
)
|
||||
assert torch.isnan(outputs).sum() == 0
|
||||
prompt_embeds = [u[:v] for u, v in zip(outputs, seq_lens)]
|
||||
prompt_embeds = torch.stack([
|
||||
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
|
||||
for u in prompt_embeds
|
||||
],
|
||||
dim=0)
|
||||
batch.prompt_embeds.append(prompt_embeds)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
negative_text = batch.negative_prompt
|
||||
negative_text_inputs = self.tokenizer(
|
||||
negative_text,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
max_length=512,
|
||||
add_special_tokens=True,
|
||||
return_attention_mask=True,
|
||||
return_tensors="pt",
|
||||
).to(batch.device)
|
||||
text_input_ids, mask = negative_text_inputs.input_ids, negative_text_inputs.attention_mask
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
negative_outputs = self.text_encoder(
|
||||
input_ids=text_input_ids,
|
||||
attention_mask=mask,
|
||||
)
|
||||
assert torch.isnan(negative_outputs).sum() == 0
|
||||
neg_prompt_embeds = [
|
||||
u[:v] for u, v in zip(negative_outputs, seq_lens)
|
||||
]
|
||||
neg_prompt_embeds = torch.stack([
|
||||
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
|
||||
for u in neg_prompt_embeds
|
||||
],
|
||||
dim=0)
|
||||
batch.negative_prompt_embeds.append(neg_prompt_embeds)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
@@ -0,0 +1,81 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Wan video diffusion pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (
|
||||
CLIPImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
|
||||
EncodingStage, InputValidationStage, LatentPreparationStage,
|
||||
T5EncodingStage, TimestepPreparationStage)
|
||||
|
||||
# TODO(will): move PRECISION_TO_TYPE to better place
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanImageToVideoPipeline(ComposedPipelineBase):
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler", \
|
||||
"image_encoder", "image_processor"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=T5EncodingStage(
|
||||
text_encoder=self.get_module("text_encoder"),
|
||||
tokenizer=self.get_module("tokenizer"),
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="image_encoding_stage",
|
||||
stage=CLIPImageEncodingStage(
|
||||
image_encoder=self.get_module("image_encoder"),
|
||||
image_processor=self.get_module("image_processor"),
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="conditioning_stage",
|
||||
stage=ConditioningStage())
|
||||
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
|
||||
inference_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
num_channels_latents = self.get_module("transformer").out_channels
|
||||
inference_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = WanImageToVideoPipeline
|
||||
@@ -0,0 +1,72 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Wan video diffusion pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
DenoisingStage, InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
T5EncodingStage,
|
||||
TimestepPreparationStage)
|
||||
|
||||
# TODO(will): move PRECISION_TO_TYPE to better place
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanPipeline(ComposedPipelineBase):
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=T5EncodingStage(
|
||||
text_encoder=self.get_module("text_encoder"),
|
||||
tokenizer=self.get_module("tokenizer"),
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="conditioning_stage",
|
||||
stage=ConditioningStage())
|
||||
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
|
||||
inference_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
num_channels_latents = self.get_module("transformer").in_channels
|
||||
inference_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = WanPipeline
|
||||
@@ -1,9 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import pytest
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.v1.distributed import (destroy_model_parallel,
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel)
|
||||
import pytest
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.v1.distributed import (init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
cleanup_dist_env_and_memory)
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
@@ -13,6 +18,9 @@ def distributed_setup():
|
||||
|
||||
This ensures proper cleanup even if tests fail.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
init_distributed_environment(world_size=1,
|
||||
rank=0,
|
||||
distributed_init_method="env://",
|
||||
@@ -23,6 +31,4 @@ def distributed_setup():
|
||||
backend="nccl")
|
||||
yield
|
||||
|
||||
if dist.is_initialized():
|
||||
destroy_model_parallel()
|
||||
dist.destroy_process_group()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# TODO: check if correct
|
||||
import os
|
||||
|
||||
@@ -19,10 +20,6 @@ logger = init_logger(__name__)
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
@@ -18,10 +19,6 @@ logger = init_logger(__name__)
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TEXT_ENCODER_PATH = os.path.join(MODEL_PATH, "text_encoder")
|
||||
TOKENIZER_PATH = os.path.join(MODEL_PATH, "tokenizer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_t5_encoder():
|
||||
# Initialize the two model implementations
|
||||
hf_config = AutoConfig.from_pretrained(TEXT_ENCODER_PATH)
|
||||
print(hf_config)
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.float16
|
||||
model1 = UMT5EncoderModel.from_pretrained(TEXT_ENCODER_PATH).to(
|
||||
precision).to(device).eval()
|
||||
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
|
||||
|
||||
loader = TextEncoderLoader()
|
||||
model2 = loader.load_model(TEXT_ENCODER_PATH, hf_config, device)
|
||||
|
||||
# Convert to float16 and move to device
|
||||
model2 = model2.to(precision)
|
||||
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("Model1 has %s parameters", len(params1))
|
||||
logger.info("Model2 has %s parameters", len(params2))
|
||||
|
||||
weight_diffs = []
|
||||
# check if embed_tokens are the same
|
||||
weights = ["encoder.block.{}.layer.0.layer_norm.weight", "encoder.block.{}.layer.0.SelfAttention.relative_attention_bias.weight", \
|
||||
"encoder.block.{}.layer.0.SelfAttention.o.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_0.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_1.weight",\
|
||||
"encoder.block.{}.layer.1.DenseReluDense.wo.weight", \
|
||||
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight", "shared.weight"]
|
||||
for idx in range(hf_config.num_hidden_layers):
|
||||
for w in weights:
|
||||
name1 = w.format(idx)
|
||||
name2 = w.format(idx)
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
assert p1.dtype == p2.dtype
|
||||
try:
|
||||
logger.info("Parameter: %s vs %s", name1, 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(" Max diff: %s, Mean diff: %s", max_diff,
|
||||
mean_diff)
|
||||
except Exception as e:
|
||||
logger.info("Error comparing %s and %s: %s", name1, name2, e)
|
||||
|
||||
# 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 T5 encoder with sample prompts")
|
||||
|
||||
with torch.no_grad():
|
||||
for prompt in prompts:
|
||||
logger.info("Testing prompt: %s", prompt)
|
||||
|
||||
# Tokenize the prompt
|
||||
tokens = tokenizer(prompt,
|
||||
padding="max_length",
|
||||
max_length=512,
|
||||
truncation=True,
|
||||
return_tensors="pt").to(device)
|
||||
|
||||
# Get outputs from HuggingFace 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).last_hidden_state
|
||||
print("--------------------------------")
|
||||
logger.info("Testing model2")
|
||||
|
||||
# Get outputs from our implementation
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs2 = model2(
|
||||
input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
)
|
||||
|
||||
# Compare last hidden states
|
||||
last_hidden_state1 = outputs1[tokens.attention_mask == 1]
|
||||
last_hidden_state2 = outputs2[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("Maximum difference in last hidden states: %s",
|
||||
max_diff_hidden.item())
|
||||
logger.info("Mean difference in last hidden states: %s",
|
||||
mean_diff_hidden.item())
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
assert mean_diff_hidden < 1e-4, \
|
||||
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
|
||||
assert max_diff_hidden < 1e-4, \
|
||||
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
|
||||
@@ -1,179 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers import AutoencoderKLHunyuanVideo as DiffusersHunyuanVAE
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.hunyuanvae import (
|
||||
AutoencoderKLHunyuanVideo as MyHunyuanVAE)
|
||||
|
||||
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()
|
||||
@@ -1,260 +0,0 @@
|
||||
import argparse
|
||||
import os
|
||||
from itertools import chain
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.distributed.device_mesh import init_device_mesh
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, destroy_distributed_environment,
|
||||
destroy_model_parallel, get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size, init_distributed_environment,
|
||||
initialize_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.dits.hunyuanvideo import (
|
||||
HunyuanVideoTransformer3DModel as HunyuanVideoDit)
|
||||
from fastvideo.v1.models.hunyuan.modules.models import (
|
||||
HYVideoDiffusionTransformer)
|
||||
from fastvideo.v1.models.loader.fsdp_load import shard_model
|
||||
from fastvideo.v1.utils.parallel_states import (
|
||||
initialize_sequence_parallel_state)
|
||||
|
||||
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,212 +0,0 @@
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.models.hunyuan.modules.models import (
|
||||
HUNYUAN_VIDEO_CONFIG, HYVideoDiffusionTransformer)
|
||||
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, destroy_distributed_environment,
|
||||
destroy_model_parallel, get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size, init_distributed_environment,
|
||||
initialize_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.dits.hunyuanvideo import (
|
||||
HunyuanVideoTransformer3DModel as HunyuanVideoDit)
|
||||
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
|
||||
|
||||
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)
|
||||
initialize_sequence_parallel_state(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") 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
|
||||
with torch.amp.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
logger.info("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}")
|
||||
|
||||
# mean diff
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
assert mean_diff < 1e-2, 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()
|
||||
@@ -1,7 +1,9 @@
|
||||
The reference videos in the `reference_videos` directory are used as part of an e2e test to ensure consistency in video generation quality across code changes. `test_inference_similarity.py` compares newly generated videos against these references using Structural Similarity Index (SSIM) metrics to detect any regressions in visual quality across code changes.
|
||||
|
||||
`reference_videos/FLASH_ATTN/` videos were generated on commit `66107fd5b8469fed25972feb632cd48887dac451`.
|
||||
`reference_videos/TORCH_SDPA/` videos were generated on commit `4ea008b8a16d7f5678a44b187ebdd7d9d0416ff1`.
|
||||
`reference_videos/FastHunyuan-diffusers/FLASH_ATTN/` videos were generated on commit `66107fd5b8469fed25972feb632cd48887dac451`.
|
||||
`reference_videos/FastHunyuan-diffusers/TORCH_SDPA/` videos were generated on commit `4ea008b8a16d7f5678a44b187ebdd7d9d0416ff1`.
|
||||
`reference_videos/Wan2.1-T2V-1.3B-Diffusers` videos were generated on commit `d085770a70988c7b26632a0c3123c24a57f7ca77`.
|
||||
`reference_videos/Wan2.1-I2V-14B-480P-Diffusers` videos were generated on commit `d085770a70988c7b26632a0c3123c24a57f7ca77`.
|
||||
|
||||
## Generation Details
|
||||
|
||||
@@ -9,7 +11,7 @@ The reference videos in the `reference_videos` directory are used as part of an
|
||||
|
||||
## Generation Parameters
|
||||
|
||||
{
|
||||
FastHunyuan-diffusers: {
|
||||
"num_gpus": 2,
|
||||
"model_path": "data/FastHunyuan-diffusers",
|
||||
"height": 720,
|
||||
@@ -26,8 +28,51 @@ The reference videos in the `reference_videos` directory are used as part of an
|
||||
"fps": 24
|
||||
}
|
||||
|
||||
### Prompts
|
||||
Wan2.1-T2V-1.3B-Diffusers: {
|
||||
"num_gpus": 2,
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 20,
|
||||
"guidance_scale": 3,
|
||||
"embedded_cfg_scale": 6,
|
||||
"flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
"text-encoder-precision": "fp32"
|
||||
}
|
||||
|
||||
1. Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
|
||||
Wan2.1-I2V-14B-480P-Diffusers: {
|
||||
"num_gpus": 2,
|
||||
"model_path": "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 6,
|
||||
"guidance_scale": 5.0,
|
||||
"embedded_cfg_scale": 6,
|
||||
"flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
"text-encoder-precision": "fp32"
|
||||
}
|
||||
|
||||
### Text-to-Video Prompts
|
||||
|
||||
1. "Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
|
||||
|
||||
2. "A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature."
|
||||
|
||||
### Image-to-Video Prompts
|
||||
|
||||
1. "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
|
||||
Image path: "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import os
|
||||
|
||||
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
|
||||
@@ -10,7 +11,7 @@ from fastvideo.v1.tests.ssim.compute_ssim import compute_video_ssim_torchvision
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Base parameters from the shell script
|
||||
BASE_PARAMS = {
|
||||
HUNYUAN_PARAMS = {
|
||||
"num_gpus": 2,
|
||||
"model_path": "FastVideo/FastHunyuan-diffusers",
|
||||
"height": 720,
|
||||
@@ -27,11 +28,66 @@ BASE_PARAMS = {
|
||||
"fps": 24,
|
||||
}
|
||||
|
||||
WAN_T2V_PARAMS = {
|
||||
"num_gpus": 2,
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 20,
|
||||
"guidance_scale": 3,
|
||||
"embedded_cfg_scale": 6,
|
||||
"flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
"text-encoder-precision": "fp32",
|
||||
}
|
||||
|
||||
WAN_I2V_PARAMS = {
|
||||
"num_gpus": 2,
|
||||
"model_path": "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 6,
|
||||
"guidance_scale": 5.0,
|
||||
"embedded_cfg_scale": 6,
|
||||
"flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
"text-encoder-precision": "fp32",
|
||||
}
|
||||
|
||||
MODEL_TO_PARAMS = {
|
||||
"FastHunyuan-diffusers": HUNYUAN_PARAMS,
|
||||
"Wan2.1-T2V-1.3B-Diffusers": WAN_T2V_PARAMS,
|
||||
}
|
||||
|
||||
I2V_MODEL_TO_PARAMS = {
|
||||
"Wan2.1-I2V-14B-480P-Diffusers": WAN_I2V_PARAMS,
|
||||
}
|
||||
|
||||
TEST_PROMPTS = [
|
||||
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.",
|
||||
"A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature."
|
||||
]
|
||||
|
||||
I2V_TEST_PROMPTS = [
|
||||
"An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot.",
|
||||
]
|
||||
|
||||
I2V_IMAGE_PATHS = [
|
||||
"https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg",
|
||||
]
|
||||
|
||||
|
||||
def write_ssim_results(output_dir, ssim_values, reference_path, generated_path,
|
||||
num_inference_steps, prompt):
|
||||
@@ -72,24 +128,28 @@ def write_ssim_results(output_dir, ssim_values, reference_path, generated_path,
|
||||
return False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("num_inference_steps", [6])
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("prompt", I2V_TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN", "TORCH_SDPA"])
|
||||
def test_inference_similarity(num_inference_steps, prompt, ATTENTION_BACKEND):
|
||||
@pytest.mark.parametrize("model_id", list(I2V_MODEL_TO_PARAMS.keys()))
|
||||
def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
"""
|
||||
Test that runs inference with different parameters and compares the output
|
||||
to reference videos using SSIM.
|
||||
"""
|
||||
assert len(I2V_TEST_PROMPTS) == len(I2V_IMAGE_PATHS), "Expect number of prompts equal to number of images"
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
|
||||
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
base_output_dir = os.path.join(script_dir, 'generated_videos')
|
||||
base_output_dir = os.path.join(script_dir, 'generated_videos', model_id)
|
||||
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
|
||||
output_video_name = f"{prompt[:100]}.mp4"
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
BASE_PARAMS = I2V_MODEL_TO_PARAMS[model_id]
|
||||
num_inference_steps = BASE_PARAMS["num_inference_steps"]
|
||||
image_path = I2V_IMAGE_PATHS[I2V_TEST_PROMPTS.index(prompt)]
|
||||
launch_args = [
|
||||
"--num-inference-steps",
|
||||
str(num_inference_steps),
|
||||
@@ -97,6 +157,8 @@ def test_inference_similarity(num_inference_steps, prompt, ATTENTION_BACKEND):
|
||||
prompt,
|
||||
"--output-path",
|
||||
output_dir,
|
||||
"--image_path",
|
||||
image_path,
|
||||
"--model-path",
|
||||
BASE_PARAMS["model_path"],
|
||||
"--height",
|
||||
@@ -123,13 +185,127 @@ def test_inference_similarity(num_inference_steps, prompt, ATTENTION_BACKEND):
|
||||
|
||||
if BASE_PARAMS["vae_sp"]:
|
||||
launch_args.append("--vae-sp")
|
||||
if "neg_prompt" in BASE_PARAMS.keys():
|
||||
launch_args.append("--neg_prompt")
|
||||
launch_args.append(BASE_PARAMS["neg_prompt"])
|
||||
if "text-encoder-precision" in BASE_PARAMS.keys():
|
||||
launch_args.append("--text-encoder-precision")
|
||||
launch_args.append(BASE_PARAMS["text-encoder-precision"])
|
||||
|
||||
launch_distributed(num_gpus=BASE_PARAMS["num_gpus"], args=launch_args)
|
||||
|
||||
assert os.path.exists(
|
||||
output_dir), f"Output video was not generated at {output_dir}"
|
||||
|
||||
reference_folder = os.path.join(script_dir, 'reference_videos', ATTENTION_BACKEND)
|
||||
reference_folder = os.path.join(script_dir, 'reference_videos', model_id, ATTENTION_BACKEND)
|
||||
|
||||
if not os.path.exists(reference_folder):
|
||||
logger.error("Reference folder missing")
|
||||
raise FileNotFoundError(
|
||||
f"Reference video folder does not exist: {reference_folder}")
|
||||
|
||||
# Find the matching reference video based on the prompt
|
||||
reference_video_name = None
|
||||
|
||||
for filename in os.listdir(reference_folder):
|
||||
if filename.endswith('.mp4') and prompt[:100] in filename:
|
||||
reference_video_name = filename
|
||||
break
|
||||
|
||||
if not reference_video_name:
|
||||
logger.error(f"Reference video not found for prompt: {prompt} with backend: {ATTENTION_BACKEND}")
|
||||
raise FileNotFoundError(f"Reference video missing")
|
||||
|
||||
reference_video_path = os.path.join(reference_folder, reference_video_name)
|
||||
generated_video_path = os.path.join(output_dir, output_video_name)
|
||||
|
||||
logger.info(
|
||||
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
|
||||
)
|
||||
ssim_values = compute_video_ssim_torchvision(reference_video_path,
|
||||
generated_video_path,
|
||||
use_ms_ssim=True)
|
||||
|
||||
mean_ssim = ssim_values[0]
|
||||
logger.info(f"SSIM mean value: {mean_ssim}")
|
||||
logger.info(f"Writing SSIM results to directory: {output_dir}")
|
||||
|
||||
success = write_ssim_results(output_dir, ssim_values, reference_video_path,
|
||||
generated_video_path, num_inference_steps,
|
||||
prompt)
|
||||
|
||||
if not success:
|
||||
logger.error("Failed to write SSIM results to file")
|
||||
|
||||
min_acceptable_ssim = 1
|
||||
assert mean_ssim >= min_acceptable_ssim, f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim}"
|
||||
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN", "TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
|
||||
def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
"""
|
||||
Test that runs inference with different parameters and compares the output
|
||||
to reference videos using SSIM.
|
||||
"""
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
|
||||
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
base_output_dir = os.path.join(script_dir, 'generated_videos', model_id)
|
||||
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
|
||||
output_video_name = f"{prompt[:100]}.mp4"
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
BASE_PARAMS = MODEL_TO_PARAMS[model_id]
|
||||
num_inference_steps = BASE_PARAMS["num_inference_steps"]
|
||||
launch_args = [
|
||||
"--num-inference-steps",
|
||||
str(num_inference_steps),
|
||||
"--prompt",
|
||||
prompt,
|
||||
"--output-path",
|
||||
output_dir,
|
||||
"--model-path",
|
||||
BASE_PARAMS["model_path"],
|
||||
"--height",
|
||||
str(BASE_PARAMS["height"]),
|
||||
"--width",
|
||||
str(BASE_PARAMS["width"]),
|
||||
"--num-frames",
|
||||
str(BASE_PARAMS["num_frames"]),
|
||||
"--guidance-scale",
|
||||
str(BASE_PARAMS["guidance_scale"]),
|
||||
"--embedded-cfg-scale",
|
||||
str(BASE_PARAMS["embedded_cfg_scale"]),
|
||||
"--flow-shift",
|
||||
str(BASE_PARAMS["flow_shift"]),
|
||||
"--seed",
|
||||
str(BASE_PARAMS["seed"]),
|
||||
"--sp-size",
|
||||
str(BASE_PARAMS["sp_size"]),
|
||||
"--tp-size",
|
||||
str(BASE_PARAMS["tp_size"]),
|
||||
"--fps",
|
||||
str(BASE_PARAMS["fps"]),
|
||||
]
|
||||
|
||||
if BASE_PARAMS["vae_sp"]:
|
||||
launch_args.append("--vae-sp")
|
||||
if "neg_prompt" in BASE_PARAMS.keys():
|
||||
launch_args.append("--neg_prompt")
|
||||
launch_args.append(BASE_PARAMS["neg_prompt"])
|
||||
if "text-encoder-precision" in BASE_PARAMS.keys():
|
||||
launch_args.append("--text-encoder-precision")
|
||||
launch_args.append(BASE_PARAMS["text-encoder-precision"])
|
||||
|
||||
launch_distributed(num_gpus=BASE_PARAMS["num_gpus"], args=launch_args)
|
||||
|
||||
assert os.path.exists(
|
||||
output_dir), f"Output video was not generated at {output_dir}"
|
||||
|
||||
reference_folder = os.path.join(script_dir, 'reference_videos', model_id, ATTENTION_BACKEND)
|
||||
|
||||
if not os.path.exists(reference_folder):
|
||||
logger.error("Reference folder missing")
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
from itertools import chain
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.dits.hunyuanvideo import (
|
||||
HunyuanVideoTransformer3DModel as HunyuanVideoDit)
|
||||
|
||||
from fastvideo.v1.models.loader.fsdp_load import shard_model
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
# Latent generated on commit c021e8a27cf437ac22827f2bc58b7f006561317f with 1 x L40S
|
||||
REFERENCE_LATENT = 1472.079828262329
|
||||
|
||||
|
||||
def initialize_identical_weights(model, seed=42):
|
||||
"""Initialize both models with identical weights using a fixed seed for reproducibility."""
|
||||
# Get all parameters from both models
|
||||
params1 = dict(model.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)
|
||||
|
||||
# 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)
|
||||
|
||||
logger.info("Model initialized with identical weights in bfloat16")
|
||||
return model
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_hunyuanvideo_distributed():
|
||||
# Get tensor parallel info
|
||||
sp_rank = get_sequence_model_parallel_rank()
|
||||
sp_world_size = get_sequence_model_parallel_world_size()
|
||||
|
||||
# Small model parameters for testing
|
||||
hidden_size = 128
|
||||
heads_num = 4
|
||||
mm_double_blocks_depth = 2
|
||||
mm_single_blocks_depth = 2
|
||||
torch.cuda.set_device("cuda:0")
|
||||
# Initialize the two model implementations
|
||||
model = 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)
|
||||
|
||||
# Initialize with identical weights
|
||||
model = initialize_identical_weights(model, seed=42)
|
||||
shard_model(model, cpu_offload=False, reshard_after_forward=True)
|
||||
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
|
||||
|
||||
model.eval()
|
||||
|
||||
# Move to GPU based on local rank (0 or 1 for 2 GPUs)
|
||||
device = torch.device(f"cuda:0")
|
||||
model = model.to(device)
|
||||
|
||||
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)
|
||||
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
output = model(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
)
|
||||
|
||||
latent = output.double().sum().item()
|
||||
|
||||
# Check if latents are similar
|
||||
diff_output_latents = abs(REFERENCE_LATENT - latent)
|
||||
logger.info(
|
||||
f"Reference latent: {REFERENCE_LATENT}, Current latent: {latent}"
|
||||
)
|
||||
assert diff_output_latents < 1e-4, f"Output latents differ significantly: max diff = {diff_output_latents}"
|
||||
@@ -0,0 +1,119 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.dits.hunyuanvideo import (
|
||||
HunyuanVideoTransformer3DModel as HunyuanVideoDit)
|
||||
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
"data", BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
CONFIG_PATH = os.path.join(TRANSFORMER_PATH, "config.json")
|
||||
|
||||
LOCAL_RANK = 0
|
||||
RANK = 0
|
||||
WORLD_SIZE = 1
|
||||
|
||||
# Latent generated on commit c021e8a27cf437ac22827f2bc58b7f006561317f with 1 x L40S
|
||||
REFERENCE_LATENT = 89.7002067565918
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_hunyuanvideo_distributed():
|
||||
logger.info(
|
||||
f"Initializing process: rank={RANK}, local_rank={LOCAL_RANK}, world_size={WORLD_SIZE}"
|
||||
)
|
||||
|
||||
torch.cuda.set_device(f"cuda:{LOCAL_RANK}")
|
||||
|
||||
# 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}"
|
||||
)
|
||||
|
||||
config = json.load(open(CONFIG_PATH))
|
||||
# remove "_class_name": "HunyuanVideoTransformer3DModel", "_diffusers_version": "0.32.0.dev0",
|
||||
# TODO: write normalize config function
|
||||
config.pop("_class_name")
|
||||
config.pop("_diffusers_version")
|
||||
|
||||
weight_dir_list = glob.glob(os.path.join(TRANSFORMER_PATH, "*.safetensors"))
|
||||
weight_dir_list = [str(path) for path in weight_dir_list]
|
||||
model = load_fsdp_model(HunyuanVideoDit,
|
||||
init_params=config,
|
||||
weight_dir_list=weight_dir_list,
|
||||
device=torch.device(f"cuda:{LOCAL_RANK}"),
|
||||
cpu_offload=False)
|
||||
|
||||
model.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)
|
||||
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
# Run inference on model
|
||||
with torch.amp.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
output = model(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
)
|
||||
|
||||
latent = output.double().sum().item()
|
||||
|
||||
# Check if latents are similar
|
||||
diff_output_latents = abs(REFERENCE_LATENT - latent)
|
||||
logger.info(
|
||||
f"Reference latent: {REFERENCE_LATENT}, Current latent: {latent}"
|
||||
)
|
||||
assert diff_output_latents < 1e-4, f"Output latents differ significantly: max diff = {diff_output_latents}"
|
||||
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from diffusers import WanTransformer3DModel
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = InferenceArgs(model_path=TRANSFORMER_PATH,
|
||||
use_cpu_offload=False,
|
||||
precision=precision_str)
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, "", args).to(device, dtype=precision)
|
||||
|
||||
model1 = WanTransformer3DModel.from_pretrained(
|
||||
TRANSFORMER_PATH, device=device,
|
||||
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# 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
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# 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("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
seq_len = 30
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([500], device=device, dtype=precision)
|
||||
|
||||
with torch.amp.autocast('cuda', dtype=precision):
|
||||
output1 = model1(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
):
|
||||
output2 = model2(hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
@@ -0,0 +1,78 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.hunyuanvae import (
|
||||
AutoencoderKLHunyuanVideo as MyHunyuanVAE)
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
"data", BASE_MODEL_PATH))
|
||||
VAE_PATH = os.path.join(MODEL_PATH, "vae")
|
||||
CONFIG_PATH = os.path.join(VAE_PATH, "config.json")
|
||||
|
||||
# Latent generated on commit 250f0b916cebb18a1c15c4aae1a0b480604d066a with 1 x A40
|
||||
REFERENCE_LATENT = -105.51324462890625
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_hunyuan_vae():
|
||||
device = torch.device("cuda:0")
|
||||
# Initialize the two model implementations
|
||||
config = json.load(open(CONFIG_PATH))
|
||||
config.pop("_class_name")
|
||||
config.pop("_diffusers_version")
|
||||
model = MyHunyuanVAE(**config).to(torch.bfloat16)
|
||||
|
||||
loaded = load_file(os.path.join(VAE_PATH,
|
||||
"diffusion_pytorch_model.safetensors"))
|
||||
model.load_state_dict(loaded)
|
||||
|
||||
# Set model to eval mode
|
||||
model.eval()
|
||||
|
||||
# Move to GPU
|
||||
model = model.to(device)
|
||||
|
||||
model.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)
|
||||
|
||||
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():
|
||||
latent = model.encode(input_tensor).mean.double().sum().item()
|
||||
|
||||
# Check if latents are similar
|
||||
diff_encoded_latents = abs(REFERENCE_LATENT - latent)
|
||||
logger.info(
|
||||
f"Reference latent: {REFERENCE_LATENT}, Current latent: {latent}"
|
||||
)
|
||||
assert diff_encoded_latents < 1e-4, f"Encoded latents differ significantly: max diff = {diff_encoded_latents}"
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from diffusers import AutoencoderKLWan
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import VAELoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
VAE_PATH = os.path.join(MODEL_PATH, "vae")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_wan_vae():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = InferenceArgs(model_path=VAE_PATH, vae_precision=precision_str)
|
||||
args.device = device
|
||||
|
||||
loader = VAELoader()
|
||||
model2 = loader.load(VAE_PATH, "", args)
|
||||
|
||||
model1 = AutoencoderKLWan.from_pretrained(
|
||||
VAE_PATH, torch_dtype=precision).to(device).eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
|
||||
# Video input [B, C, T, H, W]
|
||||
input_tensor = torch.randn(batch_size,
|
||||
3,
|
||||
81,
|
||||
32,
|
||||
32,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
latent_tensor = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
32,
|
||||
32,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
# Test encoding
|
||||
logger.info("Testing encoding...")
|
||||
latent1 = model1.encode(input_tensor).latent_dist.mean
|
||||
print("--------------------------------")
|
||||
latent2 = model2.encode(input_tensor).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("Maximum difference between encoded latents: %s",
|
||||
max_diff_encode.item())
|
||||
logger.info("Mean difference between encoded latents: %s",
|
||||
mean_diff_encode.item())
|
||||
assert mean_diff_encode < 5e-1, f"Encoded latents differ significantly: mean diff = {mean_diff_encode.item()}"
|
||||
# Test decoding
|
||||
logger.info("Testing decoding...")
|
||||
latents_mean = (torch.tensor(model1.config.latents_mean).view(
|
||||
1, model1.config.z_dim, 1, 1, 1).to(latent_tensor.device,
|
||||
latent_tensor.dtype))
|
||||
latents_std = 1.0 / torch.tensor(model1.config.latents_std).view(
|
||||
1, model1.config.z_dim, 1, 1, 1).to(latent_tensor.device,
|
||||
latent_tensor.dtype)
|
||||
latent_tensor = latent_tensor / latents_std + latents_mean
|
||||
output2 = model2.decode(latent_tensor)
|
||||
output1 = model1.decode(latent_tensor).sample
|
||||
# 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("Maximum difference between decoded outputs: %s",
|
||||
max_diff_decode.item())
|
||||
logger.info("Mean difference between decoded outputs: %s",
|
||||
mean_diff_decode.item())
|
||||
assert mean_diff_decode < 1e-1, f"Decoded outputs differ significantly: mean diff = {mean_diff_decode.item()}"
|
||||
+2
-3
@@ -19,7 +19,7 @@ dependencies = [
|
||||
|
||||
# Machine Learning & Transformers
|
||||
"transformers>=4.46.1", "tokenizers>=0.20.1", "sentencepiece==0.2.0",
|
||||
"timm==1.0.11", "peft==0.13.2", "diffusers==0.32.0", "bitsandbytes",
|
||||
"timm==1.0.11", "peft==0.13.2", "diffusers>=0.33.0", "bitsandbytes",
|
||||
"torch==2.5.1", "torchvision",
|
||||
|
||||
# vLLM
|
||||
@@ -37,9 +37,8 @@ dependencies = [
|
||||
|
||||
# Miscellaneous Utilities
|
||||
"tqdm==4.66.5", "PyYAML==6.0.1", "idna==3.6", "protobuf==5.28.3",
|
||||
"gradio==5.3.0", "huggingface_hub==0.26.1", "moviepy==1.0.3", "flask",
|
||||
"gradio==5.3.0", "moviepy==1.0.3", "flask",
|
||||
"flask_restful", "aiohttp", "huggingface_hub", "cloudpickle",
|
||||
|
||||
# System & Monitoring Tools
|
||||
"gpustat", "watch",
|
||||
|
||||
|
||||
Executable
+27
@@ -0,0 +1,27 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=2
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
export MODEL_BASE=/workspace/data/Wan2.1-T2V-1.3B-Diffusers
|
||||
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/v1/sample/v1_fastvideo_inference.py \
|
||||
--sp_size $num_gpus \
|
||||
--tp_size $num_gpus \
|
||||
--height 480 \
|
||||
--width 832 \
|
||||
--num_frames 77 \
|
||||
--num_inference_steps 50 \
|
||||
--fps 16 \
|
||||
--guidance_scale 3.0 \
|
||||
--prompt_path ./assets/prompt.txt \
|
||||
--neg_prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
|
||||
--seed 1024 \
|
||||
--output_path outputs_video/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--vae-sp \
|
||||
--text-encoder-precision "fp32" \
|
||||
--use-cpu-offload
|
||||
@@ -0,0 +1,29 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=2
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
export MODEL_BASE=/workspace/data/Wan2.1-I2V-14B-480P-Diffusers
|
||||
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/v1/sample/v1_fastvideo_inference.py \
|
||||
--sp_size $num_gpus \
|
||||
--tp_size $num_gpus \
|
||||
--height 480 \
|
||||
--width 832 \
|
||||
--num_frames 77 \
|
||||
--num_inference_steps 40 \
|
||||
--fps 16 \
|
||||
--flow_shift 3.0 \
|
||||
--guidance_scale 5.0 \
|
||||
--image_path "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg" \
|
||||
--prompt "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot." \
|
||||
--neg_prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
|
||||
--seed 1024 \
|
||||
--output_path outputs_i2v/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--vae-sp \
|
||||
--text-encoder-precision "fp32" \
|
||||
--use-cpu-offload
|
||||
Reference in New Issue
Block a user