Compare commits

..
2 Commits
Author SHA1 Message Date
William LinandJerryZhou54 008ee2099a V1 wan rebased (#335)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
2025-04-11 15:34:54 -07:00
Kevin Lin 137f61f2fe Port tests to v1 (#333) 2025-04-11 01:40:01 -07:00
63 changed files with 3547 additions and 908 deletions
+123 -5
View File
@@ -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
+5 -5
View File
@@ -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
-21
View File
@@ -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
+44
View File
@@ -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.
+22 -4
View File
@@ -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,
+2
View File
@@ -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,
+33 -2
View File
@@ -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"
)
+2
View File
@@ -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(
+53 -79
View File
@@ -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
+5
View File
@@ -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):
+35 -25
View File
@@ -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
+4 -1
View File
@@ -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
+64 -32
View 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
+10
View File
@@ -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 = [
+47
View File
@@ -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
+16 -10
View File
@@ -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,
+3
View File
@@ -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
+24 -6
View File
@@ -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)
+220
View File
@@ -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)
+15 -6
View File
@@ -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
+73 -26
View File
@@ -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.
+168
View File
@@ -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
+12 -6
View File
@@ -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()
+50 -5
View File
@@ -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
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import os
@@ -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}"
+98
View File
@@ -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
View File
@@ -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",
+27
View File
@@ -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
+29
View File
@@ -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