Compare commits

...
76 changed files with 2219 additions and 1196 deletions
+50
View File
@@ -39,6 +39,16 @@ on:
required: false
default: false
type: boolean
run_training_test:
description: "Run training-test"
required: false
default: false
type: boolean
run_nightly_test:
description: "Run nightly-test"
required: false
default: false
type: boolean
env:
PYTHONUNBUFFERED: "1"
@@ -59,6 +69,7 @@ jobs:
encoder-test: ${{ steps.filter.outputs.encoder-test }}
vae-test: ${{ steps.filter.outputs.vae-test }}
transformer-test: ${{ steps.filter.outputs.transformer-test }}
training-test: ${{ steps.filter.outputs.training-test }}
steps:
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
@@ -79,6 +90,8 @@ jobs:
- 'fastvideo/v1/tests/transformers/**'
- 'fastvideo/v1/layers/**'
- 'fastvideo/v1/attention/**'
training-test:
- 'fastvideo/v1/**'
encoder-test:
needs: change-filter
@@ -160,6 +173,43 @@ jobs:
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
training-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "training-test"
gpu_type: "NVIDIA A40"
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/training -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
nightly-test:
if: >-
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "nightly-test"
gpu_type: "NVIDIA A40"
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
runpod-cleanup:
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
+4 -1
View File
@@ -43,6 +43,8 @@ on:
required: true
RUNPOD_PRIVATE_KEY:
required: true
WANDB_API_KEY:
required: false
jobs:
run-test:
@@ -55,7 +57,7 @@ jobs:
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
python-version: "3.12"
- name: Set up SSH key
run: |
@@ -72,6 +74,7 @@ jobs:
JOB_ID: ${{ inputs.job_id }}
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
timeout-minutes: ${{ inputs.timeout_minutes }}
run: >-
python .github/scripts/runpod_api.py
+9
View File
@@ -6,6 +6,15 @@
## Installation
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
First, install C++20 for ThunderKittens:
```bash
sudo apt update
sudo apt install gcc-11 g++-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
sudo apt update
sudo apt install clang-11
```
## Environment Setup
First, set up your CUDA environment:
+1
View File
@@ -154,6 +154,7 @@ def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse
return o, lse
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
grad_output = grad_output.contiguous()
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
return grad_q, grad_k, grad_v
+15 -45
View File
@@ -7,70 +7,40 @@ To save GPU memory, we precompute text embeddings and VAE latents to eliminate t
We provide a sample dataset to help you get started. Download the source media using the following command:
```bash
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=FastVideo/mini_i2v_dataset --repo_type=dataset
```
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
To preprocess the dataset for fine-tuning or distillation, run:
```
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
bash scripts/preprocess/v1_preprocess_wan_data_t2v # for wan
```
The preprocessed dataset will be stored in `Image-Vid-Finetune-Mochi` or `Image-Vid-Finetune-HunYuan` correspondingly.
## Process your own dataset
If you wish to create your own dataset for finetuning or distillation, please structure you video dataset in the following format:
If you wish to create your own dataset for finetuning or distillation, please refer `mini_i2v_dataset/crush-smol_raw/` to structure you video dataset in the following format:
```
path_to_dataset_folder/
├── media/
│ ├── 0.jpg
path_to_your_dataset_folder/
├── videos/
│ ├── 0.mp4
│ ├── 1.mp4
│ ├── 2.jpg
├── video2caption.json
└── merge.txt
├── videos.txt
└── prompt.txt
```
Format the JSON file as a list, where each item represents a media source:
To geranate the `videos2caption.json` and `merge.txt`, run
For image media,
```
{
"path": "0.jpg",
"cap": ["captions"]
}
``` python
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
```
For video media,
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/v1_preprocess_****.sh` accordingly and run:
```
{
"path": "1.mp4",
"resolution": {
"width": 848,
"height": 480
},
"fps": 30.0,
"duration": 6.033333333333333,
"cap": [
"caption"
]
}
```
Use a txt file (merge.txt) to contain the source folder for media and the JSON file for meta information:
```
path_to_media_source_foder,path_to_json_file
```
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/preprocess_****_data.sh` accordingly and run:
```
bash scripts/preprocess/preprocess_****_data.sh
bash scripts/preprocess/v1_preprocess_****.sh
```
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
@@ -5,7 +5,11 @@ from typing import List, Optional, Type
import torch
from einops import rearrange
from vsa import video_sparse_attn
try:
from vsa import video_sparse_attn
except ImportError:
video_sparse_attn = None
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
@@ -68,14 +72,18 @@ class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
if forward_batch.latents is None:
raise ValueError("latents cannot be None")
raw_latent_shape = forward_batch.latents.shape
patch_size = fastvideo_args.dit_config.patch_size
raw_latent_shape = forward_batch.raw_latent_shape
if raw_latent_shape is None:
raise ValueError("raw_latent_shape cannot be None")
patch_size = fastvideo_args.pipeline_config.dit_config.patch_size
dit_seq_shape = [
raw_latent_shape[2] // patch_size[0],
raw_latent_shape[3] // patch_size[1],
raw_latent_shape[4] // patch_size[2]
]
VSA_sparsity = forward_batch.VSA_sparsity
return VideoSparseAttentionMetadata(current_timestep=current_timestep,
dit_seq_shape=dit_seq_shape,
VSA_sparsity=VSA_sparsity)
@@ -170,10 +178,15 @@ class VideoSparseAttentionImpl(AttentionImpl):
value = value.transpose(1, 2).contiguous()
gate_compress = gate_compress.transpose(1, 2).contiguous()
VSA_sparsity = attn_metadata.VSA_sparsity
cur_topk = math.ceil(
(1 - attn_metadata.VSA_sparsity) *
(1 - VSA_sparsity) *
(self.img_seq_length / math.prod(self.VSA_base_tile_size)))
if video_sparse_attn is None:
raise NotImplementedError("video_sparse_attn is not installed")
hidden_states = video_sparse_attn(
query,
key,
+26 -27
View File
@@ -12,7 +12,7 @@ from fastvideo.v1.distributed.communication_op import (
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
get_sp_world_size)
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
from fastvideo.v1.utils import get_compute_dtype
@@ -26,8 +26,8 @@ class DistributedAttention(nn.Module):
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
supported_attention_backends: Optional[Tuple[
AttentionBackendEnum, ...]] = None,
prefix: str = "",
**extra_impl_args) -> None:
super().__init__()
@@ -45,13 +45,13 @@ class DistributedAttention(nn.Module):
dtype,
supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.impl = impl_cls(num_heads=num_heads,
head_size=head_size,
causal=causal,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
prefix=f"{prefix}.impl",
**extra_impl_args)
self.attn_impl = impl_cls(num_heads=num_heads,
head_size=head_size,
causal=causal,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
prefix=f"{prefix}.impl",
**extra_impl_args)
self.num_heads = num_heads
self.head_size = head_size
self.num_kv_heads = num_kv_heads
@@ -100,7 +100,7 @@ class DistributedAttention(nn.Module):
scatter_dim=2,
gather_dim=1)
# Apply backend-specific preprocess_qkv
qkv = self.impl.preprocess_qkv(qkv, ctx_attn_metadata)
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
# Concatenate with replicated QKV if provided
if replicated_q is not None:
@@ -116,7 +116,7 @@ class DistributedAttention(nn.Module):
q, k, v = qkv.chunk(3, dim=0)
output = self.impl.forward(q, k, v, ctx_attn_metadata)
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
# Redistribute back if using sequence parallelism
replicated_output = None
@@ -127,7 +127,7 @@ class DistributedAttention(nn.Module):
replicated_output = sequence_model_parallel_all_gather(
replicated_output.contiguous(), dim=2)
# Apply backend-specific postprocess_output
output = self.impl.postprocess_output(output, ctx_attn_metadata)
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
@@ -183,18 +183,17 @@ class DistributedAttention_VSA(DistributedAttention):
scatter_dim=2,
gather_dim=1)
qkvg = self.impl.preprocess_qkv(
qkvg, ctx_attn_metadata) # (yongqi) pass latent shape here?
qkvg = self.attn_impl.preprocess_qkv(qkvg, ctx_attn_metadata)
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
output = self.impl.forward(q, k, v, gate_compress,
ctx_attn_metadata) # type: ignore[call-arg]
output = self.attn_impl.forward(
q, k, v, gate_compress, ctx_attn_metadata) # type: ignore[call-arg]
# Redistribute back if using sequence parallelism
replicated_output = None
# Apply backend-specific postprocess_output
output = self.impl.postprocess_output(output, ctx_attn_metadata)
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
@@ -212,8 +211,8 @@ class LocalAttention(nn.Module):
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
supported_attention_backends: Optional[Tuple[
AttentionBackendEnum, ...]] = None,
**extra_impl_args) -> None:
super().__init__()
if softmax_scale is None:
@@ -229,12 +228,12 @@ class LocalAttention(nn.Module):
dtype,
supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.impl = impl_cls(num_heads=num_heads,
head_size=head_size,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
causal=causal,
**extra_impl_args)
self.attn_impl = impl_cls(num_heads=num_heads,
head_size=head_size,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
causal=causal,
**extra_impl_args)
self.num_heads = num_heads
self.head_size = head_size
self.num_kv_heads = num_kv_heads
@@ -265,5 +264,5 @@ class LocalAttention(nn.Module):
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
output = self.impl.forward(q, k, v, ctx_attn_metadata)
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
return output
+14 -11
View File
@@ -11,13 +11,13 @@ import torch
import fastvideo.v1.envs as envs
from fastvideo.v1.attention.backends.abstract import AttentionBackend
from fastvideo.v1.logger import init_logger
from fastvideo.v1.platforms import _Backend, current_platform
from fastvideo.v1.platforms import AttentionBackendEnum, current_platform
from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
logger = init_logger(__name__)
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
def backend_name_to_enum(backend_name: str) -> Optional[AttentionBackendEnum]:
"""
Convert a string backend name to a _Backend enum value.
@@ -27,11 +27,11 @@ def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
loaded.
"""
assert backend_name is not None
return _Backend[backend_name] if backend_name in _Backend.__members__ else \
return AttentionBackendEnum[backend_name] if backend_name in AttentionBackendEnum.__members__ else \
None
def get_env_variable_attn_backend() -> Optional[_Backend]:
def get_env_variable_attn_backend() -> Optional[AttentionBackendEnum]:
'''
Get the backend override specified by the FastVideo attention
backend environment variable, if one is specified.
@@ -53,10 +53,11 @@ def get_env_variable_attn_backend() -> Optional[_Backend]:
#
# THIS SELECTION TAKES PRECEDENCE OVER THE
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
forced_attn_backend: Optional[_Backend] = None
forced_attn_backend: Optional[AttentionBackendEnum] = None
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
def global_force_attn_backend(
attn_backend: Optional[AttentionBackendEnum]) -> None:
'''
Force all attention operations to use a specified backend.
@@ -71,7 +72,7 @@ def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
forced_attn_backend = attn_backend
def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_global_forced_attn_backend() -> Optional[AttentionBackendEnum]:
'''
Get the currently-forced choice of attention backend,
or None if auto-selection is currently enabled.
@@ -82,7 +83,8 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
) -> Type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype,
supported_attention_backends)
@@ -92,7 +94,8 @@ def get_attn_backend(
def _cached_get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
) -> Type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
@@ -102,7 +105,7 @@ def _cached_get_attn_backend(
if not supported_attention_backends:
raise ValueError("supported_attention_backends is empty")
selected_backend = None
backend_by_global_setting: Optional[_Backend] = (
backend_by_global_setting: Optional[AttentionBackendEnum] = (
get_global_forced_attn_backend())
if backend_by_global_setting is not None:
selected_backend = backend_by_global_setting
@@ -125,7 +128,7 @@ def _cached_get_attn_backend(
@contextmanager
def global_force_attn_backend_context_manager(
attn_backend: _Backend) -> Generator[None, None, None]:
attn_backend: AttentionBackendEnum) -> Generator[None, None, None]:
'''
Globally force a FastVideo attention backend override within a
context manager, reverting the global attention backend
+5 -7
View File
@@ -4,7 +4,7 @@ from typing import Any, List, Optional, Tuple
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
@dataclass
@@ -13,12 +13,10 @@ class DiTArchConfig(ArchConfig):
_compile_conditions: list = field(default_factory=list)
_param_names_mapping: dict = field(default_factory=dict)
_lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA,
_Backend.VIDEO_SPARSE_ATTN)
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN)
hidden_size: int = 0
num_attention_heads: int = 0
+3 -3
View File
@@ -6,14 +6,14 @@ import torch
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
@dataclass
class EncoderArchConfig(ArchConfig):
architectures: List[str] = field(default_factory=lambda: [])
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
output_hidden_states: bool = False
use_return_dict: bool = True
+11
View File
@@ -1,4 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import dataclasses
from dataclasses import dataclass, field
from typing import Any, Union
@@ -129,3 +131,12 @@ class VAEConfig(ModelConfig):
)
return parser
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "VAEConfig":
kwargs = {}
for attr in dataclasses.fields(cls):
value = getattr(args, attr.name, None)
if value is not None:
kwargs[attr.name] = value
return cls(**kwargs)
+2 -2
View File
@@ -3,7 +3,7 @@ from fastvideo.v1.configs.pipelines.base import (PipelineConfig,
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
HunyuanConfig)
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_for_name)
get_pipeline_config_cls_from_name)
from fastvideo.v1.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
WanI2V720PConfig,
@@ -14,5 +14,5 @@ __all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"get_pipeline_config_cls_for_name"
"get_pipeline_config_cls_from_name"
]
+243 -18
View File
@@ -1,15 +1,17 @@
# SPDX-License-Identifier: Apache-2.0
import json
from dataclasses import asdict, dataclass, field, fields
from typing import Any, Callable, Dict, Optional, Tuple, cast
from typing import Any, Callable, Dict, List, Optional, Tuple, Union, cast
import torch
from fastvideo.v1.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
VAEConfig)
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput
from fastvideo.v1.configs.utils import update_config_from_args
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import shallow_asdict
from fastvideo.v1.utils import (FlexibleArgumentParser, StoreBoolean,
shallow_asdict)
logger = init_logger(__name__)
@@ -22,59 +24,282 @@ def postprocess_text(output: BaseEncoderOutput) -> torch.tensor:
raise NotImplementedError
# config for a single pipeline
@dataclass
class PipelineConfig:
"""Base configuration for all pipeline architectures."""
model_path: str = ""
pipeline_config_path: Optional[str] = None
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
dit_config: DiTConfig = field(default_factory=DiTConfig)
dit_precision: str = "bf16"
# VAE configuration
vae_config: VAEConfig = field(default_factory=VAEConfig)
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
vae_config: VAEConfig = field(default_factory=VAEConfig)
# DiT configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
# Image encoder configuration
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp16", ))
DEFAULT_TEXT_ENCODER_PRECISIONS = ("fp16", )
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), ))
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp16", ))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
...] = field(default_factory=lambda:
(postprocess_text, ))
# STA (Spatial-Temporal Attention) parameters
# LoRA parameters
lora_path: Optional[str] = None
lora_nickname: Optional[
str] = "default" # for swapping adapters in the pipeline
lora_target_names: Optional[List[
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
# StepVideo specific parameters
pos_magic: Optional[str] = None
neg_magic: Optional[str] = None
timesteps_scale: Optional[bool] = None
# STA (Sliding Tile Attention) parameters
mask_strategy_file_path: Optional[str] = None
STA_mode: str = "STA_inference"
STA_mode: Optional[str] = None
skip_time_steps: int = 15
# Compilation
enable_torch_compile: bool = False
# enable_torch_compile: bool = False
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser,
prefix: str = "") -> FlexibleArgumentParser:
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
# model_path will be conflicting with the model_path in FastVideoArgs,
# so we add it separately if prefix is not empty
if prefix_with_dot != "":
parser.add_argument(
f"--{prefix_with_dot}model-path",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}model_path",
default=PipelineConfig.model_path,
help="Path to the pretrained model",
)
parser.add_argument(
f"--{prefix_with_dot}pipeline-config-path",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}pipeline_config_path",
default=PipelineConfig.pipeline_config_path,
help="Path to the pipeline config",
)
parser.add_argument(
f"--{prefix_with_dot}embedded-cfg-scale",
type=float,
dest=f"{prefix_with_dot.replace('-', '_')}embedded_cfg_scale",
default=PipelineConfig.embedded_cfg_scale,
help="Embedded CFG scale",
)
parser.add_argument(
f"--{prefix_with_dot}flow-shift",
type=float,
dest=f"{prefix_with_dot.replace('-', '_')}flow_shift",
default=PipelineConfig.flow_shift,
help="Flow shift parameter",
)
# DiT configuration
parser.add_argument(
f"--{prefix_with_dot}dit-precision",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}dit_precision",
default=PipelineConfig.dit_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for the DiT model",
)
# VAE configuration
parser.add_argument(
f"--{prefix_with_dot}vae-precision",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}vae_precision",
default=PipelineConfig.vae_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for VAE",
)
parser.add_argument(
f"--{prefix_with_dot}vae-tiling",
action=StoreBoolean,
dest=f"{prefix_with_dot.replace('-', '_')}vae_tiling",
default=PipelineConfig.vae_tiling,
help="Enable VAE tiling",
)
parser.add_argument(
f"--{prefix_with_dot}vae-sp",
action=StoreBoolean,
dest=f"{prefix_with_dot.replace('-', '_')}vae_sp",
help="Enable VAE spatial parallelism",
)
# Text encoder configuration
parser.add_argument(
f"--{prefix_with_dot}text-encoder-precisions",
nargs="+",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}text_encoder_precisions",
default=PipelineConfig.DEFAULT_TEXT_ENCODER_PRECISIONS,
choices=["fp32", "fp16", "bf16"],
help="Precision for each text encoder",
)
# Image encoder configuration
parser.add_argument(
f"--{prefix_with_dot}image-encoder-precision",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}image_encoder_precision",
default=PipelineConfig.image_encoder_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for image encoder",
)
parser.add_argument(
f"--{prefix_with_dot}pos_magic",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}pos_magic",
default=PipelineConfig.pos_magic,
help="Positive magic prompt for sampling, used in stepvideo",
)
parser.add_argument(
f"--{prefix_with_dot}neg_magic",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}neg_magic",
default=PipelineConfig.neg_magic,
help="Negative magic prompt for sampling, used in stepvideo",
)
parser.add_argument(
f"--{prefix_with_dot}timesteps_scale",
type=bool,
dest=f"{prefix_with_dot.replace('-', '_')}timesteps_scale",
default=PipelineConfig.timesteps_scale,
help=
"Bool for applying scheduler scale in set_timesteps, used in stepvideo",
)
# Add VAE configuration arguments
from fastvideo.v1.configs.models.vaes.base import VAEConfig
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
# Add DiT configuration arguments
from fastvideo.v1.configs.models.dits.base import DiTConfig
DiTConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}dit-config")
return parser
def update_config_from_dict(self,
args: Dict[str, Any],
prefix: str = "") -> None:
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
update_config_from_args(self, args, prefix, pop_args=True)
update_config_from_args(self.vae_config,
args,
f"{prefix_with_dot}vae_config",
pop_args=True)
update_config_from_args(self.dit_config,
args,
f"{prefix_with_dot}dit_config",
pop_args=True)
@classmethod
def from_pretrained(cls, model_path: str) -> "PipelineConfig":
"""
use the pipeline class setting from model_path to match the pipeline config
"""
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_for_name)
pipeline_config_cls = get_pipeline_config_cls_for_name(model_path)
if pipeline_config_cls is not None:
pipeline_config = pipeline_config_cls()
else:
get_pipeline_config_cls_from_name)
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
@classmethod
def from_kwargs(cls,
kwargs: Dict[str, Any],
config_cli_prefix: str = "") -> "PipelineConfig":
"""
Load PipelineConfig from kwargs Dictionary.
kwargs: dictionary of kwargs
config_cli_prefix: prefix of CLI arguments for this PipelineConfig instance
"""
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
prefix_with_dot = f"{config_cli_prefix}." if (config_cli_prefix.strip()
!= "") else ""
model_path: Optional[str] = kwargs.get(prefix_with_dot + 'model_path',
None) or kwargs.get('model_path')
pipeline_config_or_path: Optional[Union[str, PipelineConfig, Dict[
str, Any]]] = kwargs.get(prefix_with_dot + 'pipeline_config',
None) or kwargs.get('pipeline_config')
if model_path is None:
raise ValueError("model_path is required in kwargs")
# 1. Get the pipeline config class from the registry
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
# 2. Instantiate PipelineConfig
if pipeline_config_cls is None:
logger.warning(
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
"Couldn't find pipeline config for %s. Using the default pipeline config.",
model_path)
pipeline_config = cls()
else:
pipeline_config = pipeline_config_cls()
return cast(PipelineConfig, pipeline_config)
# 3. Load PipelineConfig from a json file or a PipelineConfig object if provided
if isinstance(pipeline_config_or_path, str):
pipeline_config.load_from_json(pipeline_config_or_path)
kwargs[prefix_with_dot +
'pipeline_config_path'] = pipeline_config_or_path
elif isinstance(pipeline_config_or_path, PipelineConfig):
pipeline_config = pipeline_config_or_path
elif isinstance(pipeline_config_or_path, dict):
pipeline_config.update_pipeline_config(pipeline_config_or_path)
# 4. Update PipelineConfig from CLI arguments if provided
kwargs[prefix_with_dot + 'model_path'] = model_path
pipeline_config.update_config_from_dict(kwargs, config_cli_prefix)
return pipeline_config
def check_pipeline_config(self) -> None:
if self.vae_sp and not self.vae_tiling:
raise ValueError(
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
)
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
)
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
raise ValueError(
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
def dump_to_json(self, file_path: str):
output_dict = shallow_asdict(self)
+1 -1
View File
@@ -80,7 +80,7 @@ class HunyuanConfig(PipelineConfig):
(llama_postprocess_text, clip_postprocess_text))
# Precision for each component
precision: str = "bf16"
dit_precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp16"))
+63 -26
View File
@@ -19,7 +19,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
PIPE_NAME_TO_CONFIG: Dict[str, Type[PipelineConfig]] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
@@ -51,37 +51,74 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
}
def get_pipeline_config_cls_for_name(
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
"""Get the appropriate config class for specific pretrained weights."""
def get_pipeline_config_cls_from_name(
pipeline_name_or_path: str) -> Type[PipelineConfig]:
"""Get the appropriate configuration class for a given pipeline name or path.
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
logger.warning(
"FastVideo may not correctly identify the optimal config for this model, as the local directory may have been renamed."
)
else:
config = maybe_download_model_index(pipeline_name_or_path)
This function implements a multi-step lookup process to find the most suitable
configuration class for a given pipeline. It follows this order:
1. Exact match in the PIPE_NAME_TO_CONFIG
2. Partial match in the PIPE_NAME_TO_CONFIG
3. Fallback to class name in the model_index.json
4. else raise an error
pipeline_name = config["_class_name"]
Args:
pipeline_name_or_path (str): The name or path of the pipeline. This can be:
- A registered model ID (e.g., "FastVideo/FastHunyuan-diffusers")
- A local path to a model directory
- A model ID that will be downloaded
Returns:
Type[PipelineConfig]: The configuration class that best matches the pipeline.
This will be one of:
- A specific weight configuration class if an exact match is found
- A fallback configuration class based on the pipeline architecture
- The base PipelineConfig class if no matches are found
Note:
- For local paths, the function will verify the model configuration
- For remote models, it will attempt to download the model index
- Warning messages are logged when falling back to less specific configurations
"""
pipeline_config_cls: Optional[Type[PipelineConfig]] = None
# First try exact match for specific weights
if pipeline_name_or_path in WEIGHT_CONFIG_REGISTRY:
return WEIGHT_CONFIG_REGISTRY[pipeline_name_or_path]
if pipeline_name_or_path in PIPE_NAME_TO_CONFIG:
pipeline_config_cls = PIPE_NAME_TO_CONFIG[pipeline_name_or_path]
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in WEIGHT_CONFIG_REGISTRY.items():
for registered_id, config_class in PIPE_NAME_TO_CONFIG.items():
if registered_id in pipeline_name_or_path:
return config_class
# If no match, try to use the fallback config
fallback_config = None
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in PIPELINE_DETECTOR.items():
if detector(pipeline_name.lower()):
fallback_config = PIPELINE_FALLBACK_CONFIG.get(pipeline_type)
pipeline_config_cls = config_class
break
logger.warning("No match found for pipeline %s, using fallback config %s.",
pipeline_name_or_path, fallback_config)
return fallback_config
# If no match, try to use the fallback config
if pipeline_config_cls is None:
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
else:
config = maybe_download_model_index(pipeline_name_or_path)
logger.warning(
"Trying to use the config from the model_index.json. FastVideo may not correctly identify the optimal config for this model in this situation."
)
pipeline_name = config["_class_name"]
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in PIPELINE_DETECTOR.items():
if detector(pipeline_name.lower()):
pipeline_config_cls = PIPELINE_FALLBACK_CONFIG.get(
pipeline_type)
break
if pipeline_config_cls is not None:
logger.warning(
"No match found for pipeline %s, using fallback config %s.",
pipeline_name_or_path, pipeline_config_cls)
if pipeline_config_cls is None:
raise ValueError(
f"No match found for pipeline {pipeline_name_or_path}, please check the pipeline name or path."
)
return pipeline_config_cls
-7
View File
@@ -39,7 +39,6 @@ class SamplingParam:
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
VSA_sparsity: float = 0.0
# TeaCache parameters
enable_teacache: bool = False
@@ -185,12 +184,6 @@ class SamplingParam:
default=SamplingParam.image_path,
help="Path to input image for image-to-video generation",
)
parser.add_argument(
"--VSA-sparsity",
type=float,
default=SamplingParam.VSA_sparsity,
help="VSA attention sparsity",
)
return parser
+45
View File
@@ -0,0 +1,45 @@
from typing import Any, Dict
def update_config_from_args(config: Any,
args_dict: Dict[str, Any],
prefix: str = "",
pop_args: bool = False) -> None:
"""
Update configuration object from arguments dictionary.
Args:
config: The configuration object to update
args_dict: Dictionary containing arguments
prefix: Prefix for the configuration parameters in the args_dict.
If None, assumes direct attribute mapping without prefix.
"""
# Handle top-level attributes (no prefix)
args_not_to_remove = [
'model_path',
]
args_to_remove = []
if prefix.strip() == "":
for key, value in args_dict.items():
if hasattr(config, key) and value is not None:
if key == "text_encoder_precisions" and isinstance(value, list):
setattr(config, key, tuple(value))
else:
setattr(config, key, value)
if pop_args:
args_to_remove.append(key)
else:
# Handle nested attributes with prefix
prefix_with_dot = f"{prefix}."
for key, value in args_dict.items():
if key.startswith(prefix_with_dot) and value is not None:
attr_name = key[len(prefix_with_dot):]
if hasattr(config, attr_name):
setattr(config, attr_name, value)
if pop_args:
args_to_remove.append(key)
if pop_args:
for key in args_to_remove:
if key not in args_not_to_remove:
args_dict.pop(key)
@@ -0,0 +1,185 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import os
import pathlib
import time
import torch.distributed as dist
import torch.distributed.checkpoint as dist_cp
from fastvideo.v1.dataset.parquet_dataset_iterable_style import (
build_parquet_iterable_style_dataloader)
from fastvideo.v1.distributed import get_world_rank
from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_torch_device,
maybe_init_distributed_environment_and_model_parallel)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def main() -> None:
parser = argparse.ArgumentParser(
description="Benchmark parquet iterable style dataset loading speed")
parser.add_argument(
"--path",
type=str,
help="Path to parquet dataset",
)
parser.add_argument("--batch_size",
type=int,
default=4,
help="Batch size for DataLoader")
parser.add_argument("--num_data_workers",
type=int,
help="Number of DataLoader workers")
parser.add_argument("--num_epoch",
type=int,
default=2,
help="Number of epoches to benchmark")
parser.add_argument("--verify_resume",
action="store_true",
help="Verify resume")
parser.add_argument(
"--num_batches_per_epoch",
type=int,
default=1000,
help="Number of batches to benchmark",
)
parser.add_argument('--checkpoint_path',
type=str,
default='dataloader_checkpoint',
help='Path to save/load checkpoint')
'''
example launch command:
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 2 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
'''
args = parser.parse_args()
world_size = int(os.environ.get("WORLD_SIZE", 1))
maybe_init_distributed_environment_and_model_parallel(
tp_size=(world_size + 1) // 2, sp_size=(world_size + 1) // 2)
logger.info("Initialized distributed environment with world_size=%d",
world_size)
# Create DataLoader with proper settings
dataset, dataloader = build_parquet_iterable_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
logger.info("Initialized dataloader")
if args.verify_resume:
# First pass - record latent sums
first_pass_sums = []
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f", i, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
# Save dataloader state using distributed checkpoint
checkpoint_dir = pathlib.Path(args.checkpoint_path)
logger.info("Rank %d: Saving dataloader state to %s", get_world_rank(),
checkpoint_dir)
states = {"dataloader": dataloader}
begin_time = time.monotonic()
dist_cp.save(states, checkpoint_id=checkpoint_dir.as_posix())
end_time = time.monotonic()
logger.info("Rank %d: Saved checkpoint in %.2f seconds",
get_world_rank(), end_time - begin_time)
# Make sure all processes wait for checkpoint to be saved
if world_size > 1:
dist.barrier()
# Recreate dataloader and load state
dataset, dataloader = build_parquet_iterable_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
load_states = {"dataloader": dataloader}
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
logger.info("Rank %d: Loaded dataloader state from %s",
get_world_rank(), checkpoint_dir)
# Second pass - verify latent sums match
for i, (latents, embeddings, masks) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f",
i + args.num_batches_per_epoch, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
dataset, dataloader = build_parquet_iterable_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
# Second pass - verify latent sums match
second_pass_sums = []
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
second_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f (should match first pass: %f)",
i, latent_sum, first_pass_sums[i])
if i >= args.num_batches_per_epoch * 2 - 1:
break
# Verify all sums match
if all(
abs(a - b) < 1e-6
for a, b in zip(first_pass_sums, second_pass_sums)):
logger.info(
"All latent sums match between passes - resume verification successful!"
)
else:
raise ValueError(
"Latent sums do not match between passes - resume verification failed!"
)
start_time = time.time()
total_samples = 0
total_batches = 0
for _ in range(args.num_epoch):
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
if i >= args.num_batches_per_epoch:
break
# Move data to device
latents = latents.to(get_torch_device())
embeddings = embeddings.to(get_torch_device())
# Calculate actual batch size
batch_size = latents.size(0)
total_samples += batch_size
total_batches += 1
# Print progress only from rank 0
if get_world_rank() == 0 and (i + 1) % 10 == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
logger.info("Batch %d/%d, Speed: %.2f samples/sec", i + 1,
args.num_batches_per_epoch, samples_per_sec)
# Final statistics
if world_size > 1:
dist.barrier()
if get_world_rank() == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
logger.info("\nBenchmark Results:")
logger.info("Total time: %.2f seconds", elapsed)
logger.info("Total samples: %d", total_samples)
logger.info("Average speed: %.2f samples/sec", samples_per_sec)
logger.info("Time per batch: %.2f ms", elapsed / total_batches * 1000)
if __name__ == "__main__":
try:
main()
finally:
cleanup_dist_env_and_memory()
@@ -54,9 +54,9 @@ def main() -> None:
help='Path to save/load checkpoint')
'''
example launch command:
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 3 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
'''
args = parser.parse_args()
world_size = int(os.environ.get("WORLD_SIZE", 1))
@@ -66,14 +66,18 @@ def main() -> None:
world_size)
# Create DataLoader with proper settings
dataloader = build_parquet_map_style_dataloader(args.path, args.batch_size,
args.num_data_workers)
dataset, dataloader = build_parquet_map_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
logger.info("Initialized dataloader with %d batches", len(dataloader))
if args.verify_resume:
# First pass - record latent sums
first_pass_sums = []
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
logger.info("Batch %d data_indices: %s", i, data_indices)
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f", i, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
@@ -94,41 +98,54 @@ def main() -> None:
if world_size > 1:
dist.barrier()
dataloader = build_parquet_map_style_dataloader(args.path,
args.batch_size,
args.num_data_workers)
# Load dataloader state using distributed checkpoint
logger.info("Rank %d: Loading dataloader state from %s",
get_world_rank(), checkpoint_dir)
# Recreate dataloader and load state
dataset, dataloader = build_parquet_map_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
load_states = {"dataloader": dataloader}
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
logger.info("Rank %d: Loaded dataloader state from %s",
get_world_rank(), checkpoint_dir)
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
logger.info("Batch %d data_indices: %s", i, data_indices)
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f",
i + args.num_batches_per_epoch, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
logger.info("Restart from the beginning")
dataset, dataloader = build_parquet_map_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
dataloader = build_parquet_map_style_dataloader(args.path,
args.batch_size,
args.num_data_workers)
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
logger.info("Batch %d data_indices: %s", i, data_indices)
# Second pass - verify latent sums match
second_pass_sums = []
for i, (latents, embeddings, masks) in enumerate(dataloader):
latent_sum = latents.sum().item()
second_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f (should match first pass: %f)",
i, latent_sum, first_pass_sums[i])
if i >= args.num_batches_per_epoch * 2 - 1:
break
# Verify all sums match
if all(
abs(a - b) < 1e-6
for a, b in zip(first_pass_sums, second_pass_sums)):
logger.info(
"All latent sums match between passes - resume verification successful!"
)
else:
raise ValueError(
"Latent sums do not match between passes - resume verification failed!"
)
start_time = time.time()
total_samples = 0
total_batches = 0
for _ in range(args.num_epoch):
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
caption_text) in enumerate(dataloader):
if i >= args.num_batches_per_epoch:
break
-137
View File
@@ -1,137 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import json
import os
import time
from multiprocessing import Pool, cpu_count
from pathlib import Path
import torchvision
from tqdm import tqdm
def get_video_info(video_path):
"""Get video information using torchvision."""
# Read video tensor (T, C, H, W)
video_tensor, _, info = torchvision.io.read_video(str(video_path),
output_format="TCHW",
pts_unit="sec")
num_frames = video_tensor.shape[0]
height = video_tensor.shape[2]
width = video_tensor.shape[3]
fps = info.get("video_fps", 0)
duration = num_frames / fps if fps > 0 else 0
# Extract name
_, _, videos_dir, video_name = str(video_path).split("/")
return {
"path": str(video_name),
"resolution": {
"width": width,
"height": height
},
"size": os.path.getsize(video_path),
"fps": fps,
"duration": duration,
"num_frames": num_frames
}
def prepare_dataset_json(folder_path,
output_name="videos2caption.json",
num_workers=None) -> None:
"""Prepare dataset information from a folder containing videos and prompt.txt."""
folder_path = Path(folder_path)
# Read prompt file
prompt_file = folder_path / "prompt.txt"
if not prompt_file.exists():
raise FileNotFoundError(f"prompt.txt not found in {folder_path}")
with open(prompt_file) as f:
prompts = [line.strip() for line in f.readlines() if line.strip()]
# Read videos file
videos_file = folder_path / "videos.txt"
if not videos_file.exists():
raise FileNotFoundError(f"videos.txt not found in {folder_path}")
with open(videos_file) as f:
video_paths = [line.strip() for line in f.readlines() if line.strip()]
if len(prompts) != len(video_paths):
raise ValueError(
f"Number of prompts ({len(prompts)}) does not match number of videos ({len(video_paths)})"
)
# Prepare arguments for multiprocessing
process_args = [folder_path / video_path for video_path in video_paths]
# Determine number of workers
if num_workers is None:
num_workers = max(1, cpu_count() - 1) # Leave one CPU free
# Process videos in parallel
start_time = time.time()
with Pool(num_workers) as pool:
results = list(
tqdm(pool.imap(get_video_info, process_args),
total=len(process_args),
desc="Processing videos",
unit="video"))
# Combine results with prompts
dataset_info = []
for result, prompt in zip(results, prompts):
result["cap"] = [prompt]
dataset_info.append(result)
# Calculate total processing time
total_time = time.time() - start_time
total_videos = len(dataset_info)
avg_time_per_video = total_time / total_videos if total_videos > 0 else 0
print("\nProcessing completed:")
print(f"Total videos processed: {total_videos}")
print(f"Total time: {total_time:.2f} seconds")
print(f"Average time per video: {avg_time_per_video:.2f} seconds")
# Save to JSON file
output_file = folder_path / output_name
with open(output_file, 'w') as f:
json.dump(dataset_info, f, indent=2)
# Create merge.txt
merge_file = folder_path / "merge.txt"
with open(merge_file, 'w') as f:
f.write(f"{folder_path}/videos,{output_file}\n")
print(f"Dataset information saved to {output_file}")
print(f"Merge file created at {merge_file}")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description='Prepare video dataset information in JSON format')
parser.add_argument(
'--folder',
type=str,
required=True,
help='Path to the folder containing videos and prompt.txt')
parser.add_argument(
'--output',
type=str,
default='videos2caption.json',
help='Name of the output JSON file (default: videos2caption.json)')
parser.add_argument('--workers',
type=int,
default=32,
help='Number of worker processes (default: 16)')
return parser.parse_args()
if __name__ == "__main__":
args = parse_args()
prepare_dataset_json(args.folder, args.output, args.workers)
@@ -0,0 +1,275 @@
import os
import pickle
import random
from typing import Dict, List, Tuple
import numpy as np
import pyarrow.parquet as pq
import torch
import tqdm
from torch.utils.data import IterableDataset, get_worker_info
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
get_world_size)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class BatchIterator:
# TODO: Implement state_dict and load_state_dict to support resume.
def __init__(self, files, batch_size, text_padding_length, keys,
worker_num_samples, read_batch_size):
self.files = files
self.batch_size = batch_size
self.text_padding_length = text_padding_length
self.keys = keys
self.worker_num_samples = worker_num_samples
self.processed_samples = 0
self.buffer = []
self.read_batch_size = read_batch_size
def __iter__(self):
for file in self.files:
if self.processed_samples >= self.worker_num_samples:
return
reader = pq.ParquetFile(file)
for batch in reader.iter_batches(batch_size=self.read_batch_size):
if self.processed_samples >= self.worker_num_samples:
return
self.buffer.extend(batch.to_pylist())
while len(self.buffer) >= self.batch_size:
if self.processed_samples >= self.worker_num_samples:
return
batch_to_process = self.buffer[:self.batch_size]
self.buffer = self.buffer[self.batch_size:]
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
batch_to_process, self.text_padding_length, self.keys)
self.processed_samples += self.batch_size
yield all_latents, all_embs, all_masks, caption_text
class LatentsParquetIterStyleDataset(IterableDataset):
"""Efficient loader for video-text data from a directory of Parquet files."""
# Modify this in the future if we want to add more keys, for example, in image to video.
keys = [("vae_latent", "latent"), ("text_embedding")]
def __init__(self,
path: str,
batch_size: int = 1024,
cfg_rate: float = 0.1,
num_workers: int = 1,
drop_last: bool = True,
text_padding_length: int = 512,
seed: int = 42,
read_batch_size: int = 32):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.cfg_rate = cfg_rate
self.text_padding_length = text_padding_length
self.seed = seed
self.read_batch_size = read_batch_size
# Get distributed training info
self.global_rank = get_world_rank()
self.world_size = get_world_size()
self.sp_world_size = get_sp_world_size()
self.num_sp_groups = self.world_size // self.sp_world_size
num_workers = 1 if num_workers == 0 else num_workers
# Get sharding info
shard_parquet_files, shard_total_samples, shard_parquet_lengths = shard_parquet_files_across_sp_groups_and_workers(
self.path, self.num_sp_groups, num_workers, seed)
if drop_last:
self.worker_num_samples = min(
shard_total_samples) // batch_size * batch_size
# Assign files to current rank's SP group
ith_sp_group = self.global_rank // self.sp_world_size
self.sp_group_parquet_files = shard_parquet_files[ith_sp_group::self
.num_sp_groups]
self.sp_group_parquet_lengths = shard_parquet_lengths[
ith_sp_group::self.num_sp_groups]
self.sp_group_num_samples = shard_total_samples[ith_sp_group::self.
num_sp_groups]
logger.info(
"In total %d parquet files, %d samples, after sharding we retain %d samples due to drop_last",
sum([len(shard) for shard in shard_parquet_files]),
sum(shard_total_samples),
self.worker_num_samples * self.num_sp_groups * num_workers)
else:
raise ValueError("drop_last must be True")
logger.info("Each dataloader worker will load %d samples",
self.worker_num_samples)
def __iter__(self):
worker_info = get_worker_info()
worker_id = worker_info.id if worker_info is not None else 1
worker_files = self.sp_group_parquet_files[worker_id]
batch_iterator = BatchIterator(
files=worker_files,
batch_size=self.batch_size,
text_padding_length=self.text_padding_length,
keys=self.keys,
worker_num_samples=self.worker_num_samples,
read_batch_size=self.read_batch_size) # type: ignore
yield from batch_iterator
if batch_iterator.processed_samples != self.worker_num_samples:
raise ValueError(
"Rank %d, Worker %d: Not enough samples to process, this should not happen",
self.global_rank, worker_id)
def shard_parquet_files_across_sp_groups_and_workers(
path: str,
num_sp_groups: int,
num_workers: int,
seed: int = 42,
) -> Tuple[List[List[str]], List[int], List[Dict[str, int]]]:
"""
Shard parquet files across SP groups and workers in a balanced way.
Args:
path: Directory containing parquet files
num_sp_groups: Number of SP groups to shard across
num_workers: Number of workers per SP group
seed: Random seed for shuffling
Returns:
Tuple containing:
- List of lists of parquet files for each shard
- List of total samples per shard
- List of dictionaries mapping file paths to their lengths
"""
# Check if sharding plan already exists
sharding_info_dir = os.path.join(
path, f"sharding_info_{num_sp_groups}_sp_groups_{num_workers}_workers")
if os.path.exists(sharding_info_dir):
logger.info("Sharding plan already exists")
logger.info("Loading sharding plan from %s", sharding_info_dir)
try:
with open(
os.path.join(sharding_info_dir, "shard_parquet_files.pkl"),
"rb") as f:
shard_parquet_files = pickle.load(f)
with open(
os.path.join(sharding_info_dir, "shard_total_samples.pkl"),
"rb") as f:
shard_total_samples = pickle.load(f)
with open(
os.path.join(sharding_info_dir,
"shard_parquet_lengths.pkl"), "rb") as f:
shard_parquet_lengths = pickle.load(f)
return shard_parquet_files, shard_total_samples, shard_parquet_lengths
except Exception as e:
logger.error("Error loading sharding plan: %s", str(e))
logger.info("Falling back to creating new sharding plan")
if get_world_rank() == 0:
logger.info("Scanning for parquet files in %s", path)
# Find all parquet files
parquet_files = []
for root, _, files in os.walk(path):
for file in files:
if file.endswith('.parquet'):
parquet_files.append(os.path.join(root, file))
if not parquet_files:
raise ValueError("No parquet files found in %s", path)
# Calculate file lengths efficiently using a single pass
logger.info("Calculating file lengths...")
lengths = []
for file in tqdm.tqdm(parquet_files, desc="Reading parquet files"):
lengths.append(pq.ParquetFile(file).metadata.num_rows)
total_samples = sum(lengths)
logger.info("Found %d files with %d total samples", len(parquet_files),
total_samples)
# Sort files by length for better balancing
sorted_indices = np.argsort(lengths)
sorted_files = [parquet_files[i] for i in sorted_indices]
sorted_lengths = [lengths[i] for i in sorted_indices]
# Create shards
num_shards = num_sp_groups * num_workers
shard_parquet_files = [[] for _ in range(num_shards)]
shard_total_samples = [0] * num_shards
shard_parquet_lengths = [{} for _ in range(num_shards)]
# Distribute files to shards using a greedy approach
logger.info("Distributing files to shards...")
for file, length in zip(reversed(sorted_files),
reversed(sorted_lengths)):
# Find shard with minimum current length
target_shard = np.argmin(shard_total_samples)
shard_parquet_files[target_shard].append(file)
shard_total_samples[target_shard] += length
shard_parquet_lengths[target_shard][file] = length
#randomize each shard
for shard in shard_parquet_files:
random.seed(seed)
random.shuffle(shard)
save_dir = os.path.join(
path,
f"sharding_info_{num_sp_groups}_sp_groups_{num_workers}_workers")
os.makedirs(save_dir, exist_ok=True)
with open(os.path.join(save_dir, "shard_parquet_files.pkl"), "wb") as f:
pickle.dump(shard_parquet_files, f)
with open(os.path.join(save_dir, "shard_total_samples.pkl"), "wb") as f:
pickle.dump(shard_total_samples, f)
with open(os.path.join(save_dir, "shard_parquet_lengths.pkl"),
"wb") as f:
pickle.dump(shard_parquet_lengths, f)
logger.info("Saved sharding info to %s", save_dir)
# wait for all ranks to finish
torch.distributed.barrier()
# recursive call
return shard_parquet_files_across_sp_groups_and_workers(
path, num_sp_groups, num_workers, seed)
def build_parquet_iterable_style_dataloader(
path: str,
batch_size: int,
num_data_workers: int,
cfg_rate: float = 0.0,
drop_last: bool = True,
text_padding_length: int = 512,
seed: int = 42,
read_batch_size: int = 32
) -> Tuple[LatentsParquetIterStyleDataset, StatefulDataLoader]:
"""Build a dataloader for the LatentsParquetIterStyleDataset."""
dataset = LatentsParquetIterStyleDataset(
path=path,
batch_size=batch_size,
cfg_rate=cfg_rate,
num_workers=num_data_workers,
drop_last=drop_last,
text_padding_length=text_padding_length,
seed=seed,
read_batch_size=read_batch_size)
loader = StatefulDataLoader(
dataset,
batch_size=1,
num_workers=num_data_workers,
pin_memory=True,
)
return dataset, loader
@@ -1,15 +1,17 @@
# SPDX-License-Identifier: Apache-2.0
import os
import pickle
from typing import Any, Dict, List, Tuple
import numpy as np
import pyarrow.parquet as pq
# Torch in general
import torch
import tqdm
# Dataset
from torch.utils.data import Dataset, Sampler
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
get_world_size)
from fastvideo.v1.logger import init_logger
@@ -30,6 +32,7 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
sp_world_size: int,
global_rank: int,
drop_last: bool = True,
drop_first_row: bool = False,
seed: int = 0,
):
self.batch_size = batch_size
@@ -45,6 +48,11 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
# Create a random permutation of all indices
global_indices = torch.randperm(self.dataset_size, generator=rng)
if drop_first_row:
# drop 0 in global_indices
global_indices = global_indices[global_indices != 0]
self.dataset_size = self.dataset_size - 1
if self.drop_last:
# For drop_last=True, we:
# 1. Ensure total samples is divisible by (batch_size * num_sp_groups)
@@ -56,19 +64,22 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
self.num_sp_groups *
self.batch_size]
else:
# add more indices to make it divisible by (batch_size * num_sp_groups)
padding_size = self.num_sp_groups * self.batch_size - (
self.dataset_size % (self.num_sp_groups * self.batch_size))
global_indices = torch.cat(
[global_indices, global_indices[:padding_size]])
if self.dataset_size % (self.num_sp_groups * self.batch_size) != 0:
# add more indices to make it divisible by (batch_size * num_sp_groups)
padding_size = self.num_sp_groups * self.batch_size - (
self.dataset_size % (self.num_sp_groups * self.batch_size))
logger.info("Padding the dataset from %d to %d",
self.dataset_size, self.dataset_size + padding_size)
global_indices = torch.cat(
[global_indices, global_indices[:padding_size]])
# shard the indices to each sp group
ith_sp_group = self.global_rank // self.sp_world_size
sp_group_local_indices = global_indices[ith_sp_group::self.
num_sp_groups]
self.sp_group_local_indices = sp_group_local_indices
logger.info("sp_group_local_indices: %d", len(sp_group_local_indices))
logger.info("Dataset size for each sp group: %d",
len(sp_group_local_indices))
def __iter__(self):
indices = self.sp_group_local_indices
@@ -81,20 +92,49 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
def get_parquet_files_and_length(path: str):
lengths = []
file_names = []
for root, _, files in os.walk(path):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.join(root, file)
num_rows = pq.ParquetFile(file_path).metadata.num_rows
lengths.append(num_rows)
file_names.append(file_path)
# sort according to file name to ensure all rank has the same order (in case os.walk is not sorted)
file_names_sorted, lengths_sorted = zip(
*sorted(zip(file_names, lengths), key=lambda x: x[0]))
assert len(file_names_sorted) != 0, "No parquet files found in the dataset"
return file_names_sorted, lengths_sorted
# Check if cached info exists
cache_dir = os.path.join(path, "map_style_cache")
cache_file = os.path.join(cache_dir, "file_info.pkl")
if os.path.exists(cache_file):
logger.info("Loading cached file info from %s", cache_file)
try:
with open(cache_file, "rb") as f:
file_names_sorted, lengths_sorted = pickle.load(f)
return file_names_sorted, lengths_sorted
except Exception as e:
logger.error("Error loading cached file info: %s", str(e))
logger.info("Falling back to scanning files")
# If no cache exists or loading failed, scan files
if get_world_rank() == 0:
lengths = []
file_names = []
for root, _, files in os.walk(path):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.join(root, file)
file_names.append(file_path)
for file_path in tqdm.tqdm(file_names,
desc="Reading parquet files to get lengths"):
num_rows = pq.ParquetFile(file_path).metadata.num_rows
lengths.append(num_rows)
# sort according to file name to ensure all rank has the same order (in case os.walk is not sorted)
file_names_sorted, lengths_sorted = zip(
*sorted(zip(file_names, lengths), key=lambda x: x[0]))
assert len(
file_names_sorted) != 0, "No parquet files found in the dataset"
os.makedirs(cache_dir, exist_ok=True)
with open(cache_file, "wb") as f:
pickle.dump((file_names_sorted, lengths_sorted), f)
logger.info("Saved file info to %s", cache_file)
# Wait for rank 0 to finish saving
if get_world_size() > 1:
torch.distributed.barrier()
return get_parquet_files_and_length(path)
def read_row_from_parquet_file(parquet_files: List[str], global_row_idx: int,
@@ -145,7 +185,7 @@ class LatentsParquetMapStyleDataset(Dataset):
Using parquet for map style dataset is not efficient, we mainly keep it for backward compatibility and debugging.
"""
# Modify this in the future if we want to add more keys, for example, in image to video.
keys = ["vae_latent", "text_embedding"]
keys = [("vae_latent", "latent"), "text_embedding"]
def __init__(
self,
@@ -154,6 +194,7 @@ class LatentsParquetMapStyleDataset(Dataset):
cfg_rate: float = 0.0,
seed: int = 42,
drop_last: bool = True,
drop_first_row: bool = False,
text_padding_length: int = 512,
):
super().__init__()
@@ -184,27 +225,14 @@ class LatentsParquetMapStyleDataset(Dataset):
sp_world_size=get_sp_world_size(),
global_rank=get_world_rank(),
drop_last=drop_last,
drop_first_row=drop_first_row,
seed=seed,
)
logger.info("Dataset initialized with %d parquet files and %d rows",
len(self.parquet_files), sum(self.lengths))
def _get_torch_tensors_from_row_dict(
self, row_dict: Dict[str, Any]) -> Dict[str, torch.Tensor]:
"""
Get the latents and prompts from a row dictionary.
"""
return_dict = {}
for key in self.keys:
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
# TODO (peiyuan): read precision
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
return_dict[key] = data
return return_dict
def get_validation_negative_prompt(self) -> tuple[Any, Any, Any, Any]:
def get_validation_negative_prompt(
self) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, str]:
"""
Get the negative prompt for validation.
This method ensures the negative prompt is loaded and cached properly.
@@ -218,36 +246,16 @@ class LatentsParquetMapStyleDataset(Dataset):
row_dict = read_row_from_parquet_file([file_path], row_idx,
[self.lengths[0]])
# Get tensors using the existing helper method
data = self._get_torch_tensors_from_row_dict(row_dict)
emb = data["text_embedding"]
# Pad the embedding and get mask
padded_emb, mask = self._pad(emb, self.text_padding_length)
# Pin memory for faster transfer to GPU
padded_emb = padded_emb
mask = mask
return None, padded_emb, mask, None
def _pad(self, t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
Pad or crop an embedding [L, D] to exactly padding_length tokens.
Return:
- [L, D] tensor in pinned CPU memory
- [L] attention mask in pinned CPU memory
"""
L, D = t.shape
if padding_length > L: # pad
pad = torch.zeros(padding_length - L,
D,
dtype=t.dtype,
device=t.device)
return torch.cat([t, pad], 0), torch.cat(
[torch.ones(L), torch.zeros(padding_length - L)], 0)
else: # crop
return t[:padding_length], torch.ones(padding_length)
all_latents_list, all_embs_list, all_masks_list, caption_text_list = collate_latents_embs_masks(
[row_dict], self.text_padding_length, self.keys)
all_latents, all_embs, all_masks, caption_text = all_latents_list[
0], all_embs_list[0], all_masks_list[0], caption_text_list[0]
# add batch dimension
if len(all_embs.shape) == 2:
all_embs = all_embs.unsqueeze(0)
if len(all_masks.shape) == 1:
all_masks = all_masks.unsqueeze(0).unsqueeze(0)
return all_latents, all_embs, all_masks, caption_text
# PyTorch calls this ONLY because the batch_sampler yields a list
def __getitems__(self, indices: List[int]):
@@ -259,29 +267,9 @@ class LatentsParquetMapStyleDataset(Dataset):
for idx in indices
]
# Initialize tensors to hold padded embeddings and masks
all_latents = []
all_embs = []
all_masks = []
# Process each row individually
for i, row in enumerate(rows):
# Get tensors from row
data = self._get_torch_tensors_from_row_dict(row)
latents, emb = data["vae_latent"], data["text_embedding"]
padded_emb, mask = self._pad(emb, self.text_padding_length)
# Store in batch tensors
all_latents.append(latents)
all_embs.append(padded_emb)
all_masks.append(mask)
# Pin memory for faster transfer to GPU
all_latents = torch.stack(all_latents)
all_embs = torch.stack(all_embs)
all_masks = torch.stack(all_masks)
return all_latents, all_embs, all_masks, indices
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
rows, self.text_padding_length, self.keys)
return all_latents, all_embs, all_masks, caption_text
def __len__(self):
return sum(self.lengths)
@@ -300,6 +288,7 @@ def build_parquet_map_style_dataloader(
num_data_workers,
cfg_rate=0.0,
drop_last=True,
drop_first_row=False,
text_padding_length=512,
seed=42) -> Tuple[LatentsParquetMapStyleDataset, StatefulDataLoader]:
dataset = LatentsParquetMapStyleDataset(
@@ -307,6 +296,7 @@ def build_parquet_map_style_dataloader(
batch_size,
cfg_rate=cfg_rate,
drop_last=drop_last,
drop_first_row=drop_first_row,
text_padding_length=text_padding_length,
seed=seed)
+85
View File
@@ -0,0 +1,85 @@
from typing import Any, Dict, List
import numpy as np
import torch
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
Pad or crop an embedding [L, D] to exactly padding_length tokens.
Return:
- [L, D] tensor in pinned CPU memory
- [L] attention mask in pinned CPU memory
"""
L, D = t.shape
if padding_length > L: # pad
pad = torch.zeros(padding_length - L, D, dtype=t.dtype, device=t.device)
return torch.cat([t, pad], 0), torch.cat(
[torch.ones(L), torch.zeros(padding_length - L)], 0)
else: # crop
return t[:padding_length], torch.ones(padding_length)
def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
"""
Get the latents and prompts from a row dictionary.
"""
return_dict = {}
for key in keys:
shape, bytes = None, None
if isinstance(key, tuple):
for k in key:
try:
shape = row_dict[f"{k}_shape"]
bytes = row_dict[f"{k}_bytes"]
except KeyError:
continue
key = key[0]
if shape is None or bytes is None:
raise ValueError(f"Key {key} not found in row_dict")
else:
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
# TODO (peiyuan): read precision
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
if len(data.shape) == 3:
B, L, D = data.shape
assert B == 1, "Batch size must be 1"
data = data.squeeze(0)
return_dict[key] = data
return return_dict
def collate_latents_embs_masks(
batch_to_process, text_padding_length,
keys) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str]]:
# Initialize tensors to hold padded embeddings and masks
all_latents = []
all_embs = []
all_masks = []
caption_text = []
# Process each row individually
for i, row in enumerate(batch_to_process):
# Get tensors from row
data = get_torch_tensors_from_row_dict(row, keys)
latents, emb = data["vae_latent"], data["text_embedding"]
padded_emb, mask = pad(emb, text_padding_length)
# Store in batch tensors
all_latents.append(latents)
all_embs.append(padded_emb)
all_masks.append(mask)
# TODO(py): remove this once we fix preprocess
try:
caption_text.append(row["prompt"])
except KeyError:
caption_text.append(row["caption"])
# Pin memory for faster transfer to GPU
all_latents = torch.stack(all_latents)
all_embs = torch.stack(all_embs)
all_masks = torch.stack(all_masks)
return all_latents, all_embs, all_masks, caption_text
+6 -47
View File
@@ -4,9 +4,9 @@
import argparse
import dataclasses
import os
from typing import Any, Dict, List, Optional, cast
from typing import List, cast
from fastvideo import PipelineConfig, VideoGenerator
from fastvideo import VideoGenerator
from fastvideo.v1.configs.sample.base import SamplingParam
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.v1.entrypoints.cli.utils import RaiseNotImplementedAction
@@ -37,8 +37,6 @@ class GenerateSubcommand(CLISubcommand):
def cmd(self, args: argparse.Namespace) -> None:
excluded_args = ['subparser', 'config', 'dispatch_function']
FastVideoArgs.from_cli_args(args)
provided_args = {}
for k, v in vars(args).items():
if (k not in excluded_args and v is not None
@@ -66,27 +64,19 @@ class GenerateSubcommand(CLISubcommand):
init_args = {
k: v
for k, v in merged_args.items() if k in self.init_arg_names
for k, v in merged_args.items()
if k not in self.generation_arg_names
}
generation_args = {
k: v
for k, v in merged_args.items() if k in self.generation_arg_names
}
pipeline_config = PipelineConfig.from_pretrained(
merged_args['model_path'])
update_config_from_args(pipeline_config.dit_config, merged_args,
"dit_config")
update_config_from_args(pipeline_config.vae_config, merged_args,
"vae_config")
update_config_from_args(pipeline_config, merged_args)
model_path = init_args.pop('model_path')
prompt = generation_args.pop('prompt')
generator = VideoGenerator.from_pretrained(
model_path=model_path, **init_args, pipeline_config=pipeline_config)
generator = VideoGenerator.from_pretrained(model_path=model_path,
**init_args)
generator.generate_video(prompt=prompt, **generation_args)
@@ -132,34 +122,3 @@ class GenerateSubcommand(CLISubcommand):
def cmd_init() -> List[CLISubcommand]:
return [GenerateSubcommand()]
def update_config_from_args(config: Any,
args_dict: Dict[str, Any],
prefix: Optional[str] = None) -> None:
"""
Update configuration object from arguments dictionary.
Args:
config: The configuration object to update
args_dict: Dictionary containing arguments
prefix: Prefix for the configuration parameters in the args_dict.
If None, assumes direct attribute mapping without prefix.
"""
# Handle top-level attributes (no prefix)
if prefix is None:
for key, value in args_dict.items():
if hasattr(config, key) and value is not None:
if key == "text_encoder_precisions" and isinstance(value, list):
setattr(config, key, tuple(value))
else:
setattr(config, key, value)
return
# Handle nested attributes with prefix
prefix_with_dot = f"{prefix}."
for key, value in args_dict.items():
if key.startswith(prefix_with_dot) and value is not None:
attr_name = key[len(prefix_with_dot):]
if hasattr(config, attr_name):
setattr(config, attr_name, value)
+12 -34
View File
@@ -18,8 +18,6 @@ import torch
import torchvision
from einops import rearrange
from fastvideo.v1.configs.pipelines import (PipelineConfig,
get_pipeline_config_cls_for_name)
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -55,9 +53,6 @@ class VideoGenerator:
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[
Union[str
| PipelineConfig]] = None,
**kwargs) -> "VideoGenerator":
"""
Create a video generator from a pretrained model.
@@ -66,35 +61,17 @@ class VideoGenerator:
model_path: Path or identifier for the pretrained model
device: Device to load the model on (e.g., "cuda", "cuda:0", "cpu")
torch_dtype: Data type for model weights (e.g., torch.float16)
**kwargs: Additional arguments to customize model loading
pipeline_config: Pipeline config to use for inference
**kwargs: Additional arguments to customize model loading, set any FastVideoArgs or PipelineConfig attributes here.
Returns:
The created video generator
Priority level: Default pipeline config < User's pipeline config < User's kwargs
"""
config = None
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, PipelineConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = kwargs
else:
config_args = shallow_asdict(config)
config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=model_path, **config_args)
# If users also provide some kwargs, it will override the FastVideoArgs and PipelineConfig.
kwargs['model_path'] = model_path
fastvideo_args = FastVideoArgs.from_kwargs(kwargs)
return cls.from_fastvideo_args(fastvideo_args)
@@ -150,16 +127,17 @@ class VideoGenerator:
"""
# Create a copy of inference args to avoid modifying the original
fastvideo_args = self.fastvideo_args
pipeline_config = fastvideo_args.pipeline_config
# Validate inputs
if not isinstance(prompt, str):
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
kwargs["prompt"] = prompt
sampling_param.update(kwargs)
@@ -176,10 +154,10 @@ class VideoGenerator:
f"height={sampling_param.height}, width={sampling_param.width}, "
f"num_frames={sampling_param.num_frames}")
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
temporal_scale_factor = pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = sampling_param.num_frames
num_gpus = fastvideo_args.num_gpus
use_temporal_scaling_frames = fastvideo_args.vae_config.use_temporal_scaling_frames
use_temporal_scaling_frames = pipeline_config.vae_config.use_temporal_scaling_frames
# Adjust number of frames based on number of GPUs
if use_temporal_scaling_frames:
@@ -238,18 +216,18 @@ class VideoGenerator:
num_videos_per_prompt: {sampling_param.num_videos_per_prompt}
guidance_scale: {sampling_param.guidance_scale}
n_tokens: {n_tokens}
flow_shift: {fastvideo_args.flow_shift}
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}
flow_shift: {fastvideo_args.pipeline_config.flow_shift}
embedded_guidance_scale: {fastvideo_args.pipeline_config.embedded_cfg_scale}
save_video: {sampling_param.save_video}
output_path: {sampling_param.output_path}
""" # type: ignore[attr-defined]
logger.info(debug_str)
# Prepare batch
batch = ForwardBatch(
**shallow_asdict(sampling_param),
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
extra={},
)
+72 -201
View File
@@ -6,26 +6,32 @@ import argparse
import dataclasses
from contextlib import contextmanager
from dataclasses import field
from typing import Any, Callable, List, Optional, Tuple
from typing import Any, Dict, List, Optional
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import FlexibleArgumentParser, StoreBoolean
logger = init_logger(__name__)
def preprocess_text(prompt: str) -> str:
return prompt
def clean_cli_args(args: argparse.Namespace) -> Dict[str, Any]:
"""
Clean the arguments by removing the ones that not explicitly provided by the user.
"""
provided_args = {}
for k, v in vars(args).items():
if (v is not None and hasattr(args, '_provided')
and k in args._provided):
provided_args[k] = v
def postprocess_text(output: Any) -> Any:
raise NotImplementedError
return provided_args
# args for fastvideo framework
@dataclasses.dataclass
class FastVideoArgs:
# Model and path configuration
# Model and path configuration (for convenience)
model_path: str
# Cache strategy
@@ -48,66 +54,25 @@ class FastVideoArgs:
hsdp_shard_dim: int = -1
dist_timeout: Optional[int] = None # timeout for torch.distributed
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
pipeline_config: PipelineConfig = field(default_factory=PipelineConfig)
output_type: str = "pil"
# DiT configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
precision: str = "bf16"
use_cpu_offload: bool = True
use_fsdp_inference: bool = True
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True # Might change in between forward passes
vae_sp: bool = False # Might change in between forward passes
# vae_scale_factor: Optional[int] = None # Deprecated
vae_config: VAEConfig = field(default_factory=VAEConfig)
# Image encoder configuration
image_encoder_precision: str = "fp32"
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
# Text encoder configuration
DEFAULT_TEXT_ENCODER_PRECISIONS = (
"fp16",
# "fp16",
)
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), ))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: Tuple[Callable[[Any], Any], ...] = field(
default_factory=lambda: (postprocess_text, ))
# STA parameters
# STA (Sliding Tile Attention) parameters
mask_strategy_file_path: Optional[str] = None
STA_mode: Optional[str] = None
skip_time_steps: int = 15
# LoRA parameters
lora_path: Optional[str] = None
lora_nickname: Optional[
str] = "default" # for swapping adapters in the pipeline
lora_target_names: Optional[List[
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
# STA parameters
mask_strategy_file_path: Optional[str] = None
# Compilation
enable_torch_compile: bool = False
disable_autocast: bool = False
# StepVideo specific parameters
pos_magic: Optional[str] = None
neg_magic: Optional[str] = None
timesteps_scale: Optional[bool] = None
# Logging
log_level: str = "info"
# VSA parameters
VSA_sparsity: float = 0.0 # inference/validation sparsity
@property
def training_mode(self) -> bool:
@@ -125,11 +90,6 @@ class FastVideoArgs:
help=
"The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
)
parser.add_argument(
"--dit-weight",
type=str,
help="Path to the DiT model weights",
)
parser.add_argument(
"--model-dir",
type=str,
@@ -175,14 +135,12 @@ class FastVideoArgs:
help="The number of GPUs to use.",
)
parser.add_argument(
"--tensor-parallel-size",
"--tp-size",
type=int,
default=FastVideoArgs.tp_size,
help="The tensor parallelism size.",
)
parser.add_argument(
"--sequence-parallel-size",
"--sp-size",
type=int,
default=FastVideoArgs.sp_size,
@@ -207,19 +165,7 @@ class FastVideoArgs:
help="Set timeout for torch.distributed initialization.",
)
parser.add_argument(
"--embedded-cfg-scale",
type=float,
default=FastVideoArgs.embedded_cfg_scale,
help="Embedded CFG scale",
)
parser.add_argument(
"--flow-shift",
"--shift",
type=float,
default=FastVideoArgs.flow_shift,
help="Flow shift parameter",
)
# Output type
parser.add_argument(
"--output-type",
type=str,
@@ -228,53 +174,7 @@ class FastVideoArgs:
help="Output type for the generated video",
)
parser.add_argument(
"--precision",
type=str,
default=FastVideoArgs.precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for the model",
)
# VAE configuration
parser.add_argument(
"--vae-precision",
type=str,
default=FastVideoArgs.vae_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for VAE",
)
parser.add_argument(
"--vae-tiling",
action=StoreBoolean,
default=FastVideoArgs.vae_tiling,
help="Enable VAE tiling",
)
parser.add_argument(
"--vae-sp",
action=StoreBoolean,
help="Enable VAE spatial parallelism",
)
parser.add_argument(
"--text-encoder-precisions",
nargs="+",
type=str,
default=FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS,
choices=["fp32", "fp16", "bf16"],
help="Precision for each text encoder",
)
# Image encoder config
parser.add_argument(
"--image-encoder-precision",
type=str,
default=FastVideoArgs.image_encoder_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for image encoder",
)
# STA parameters
# STA (Sliding Tile Attention) parameters
parser.add_argument(
"--STA-mode",
type=str,
@@ -323,69 +223,42 @@ class FastVideoArgs:
"Disable autocast for denoising loop and vae decoding in pipeline sampling",
)
# VSA parameters
parser.add_argument(
"--pos_magic",
type=str,
default=FastVideoArgs.pos_magic,
help="Positive magic prompt for sampling",
)
parser.add_argument(
"--neg_magic",
type=str,
default=FastVideoArgs.neg_magic,
help="Negative magic prompt for sampling",
)
parser.add_argument(
"--timesteps_scale",
type=bool,
default=FastVideoArgs.timesteps_scale,
help="Bool for applying scheduler scale in set_timesteps",
"--VSA-sparsity",
type=float,
default=FastVideoArgs.VSA_sparsity,
help="Validation sparsity for VSA",
)
# Logging
parser.add_argument(
"--log-level",
type=str,
default=FastVideoArgs.log_level,
help="The logging level of all loggers.",
)
# Add VAE configuration arguments
from fastvideo.v1.configs.models.vaes.base import VAEConfig
VAEConfig.add_cli_args(parser)
# Add DiT configuration arguments
from fastvideo.v1.configs.models.dits.base import DiTConfig
DiTConfig.add_cli_args(parser)
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
return parser
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs":
args.tp_size = args.tensor_parallel_size
args.sp_size = args.sequence_parallel_size
args.flow_shift = getattr(args, "shift", args.flow_shift)
provided_args = clean_cli_args(args)
# Get all fields from the dataclass
attrs = [attr.name for attr in dataclasses.fields(cls)]
# Create a dictionary of attribute values, with defaults for missing attributes
kwargs = {}
for attr in attrs:
# Handle renamed attributes or those with multiple CLI names
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
if attr == 'pipeline_config':
pipeline_config = PipelineConfig.from_kwargs(provided_args)
kwargs[attr] = pipeline_config
# Use getattr with default value from the dataclass for potentially missing attributes
else:
default_value = getattr(cls, attr, None)
value = getattr(args, attr, default_value)
if value is not None:
kwargs[attr] = value
kwargs[attr] = value # type: ignore
return cls(**kwargs) # type: ignore
@classmethod
def from_kwargs(cls, kwargs: Dict[str, Any]) -> "FastVideoArgs":
kwargs['pipeline_config'] = PipelineConfig.from_kwargs(kwargs)
return cls(**kwargs)
def check_fastvideo_args(self) -> None:
@@ -414,33 +287,17 @@ class FastVideoArgs:
f"tp_size ({self.tp_size}) must be equal to sp_size ({self.sp_size})"
)
# Validate VAE spatial parallelism with VAE tiling
if self.vae_sp and not self.vae_tiling:
raise ValueError(
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
)
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
)
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
raise ValueError(
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
if self.enable_torch_compile and self.num_gpus > 1:
logger.warning(
"Currently torch compile does not work with multi-gpu. Setting enable_torch_compile to False"
)
self.enable_torch_compile = False
if self.pipeline_config is None:
raise ValueError("pipeline_config is not set in FastVideoArgs")
self.pipeline_config.check_pipeline_config()
_current_fastvideo_args = None
@@ -514,7 +371,6 @@ class TrainingArgs(FastVideoArgs):
# text encoder & vae & diffusion model
pretrained_model_name_or_path: str = ""
dit_model_name_or_path: str = ""
cache_dir: str = ""
# diffusion setting
ema_decay: float = 0.0
@@ -529,6 +385,7 @@ class TrainingArgs(FastVideoArgs):
validation_steps: float = 0.0
log_validation: bool = False
tracker_project_name: str = ""
wandb_run_name: str = ""
seed: Optional[int] = None
# output
@@ -576,30 +433,29 @@ class TrainingArgs(FastVideoArgs):
# master_weight_type
master_weight_type: str = ""
# For fast checking in LoRA pipeline
training_mode: bool = True
# VSA training decay parameters
VSA_decay_rate: float = 0.01 # decay rate -> 0.02
VSA_decay_interval_steps: int = 1 # decay interval steps -> 50
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
provided_args = clean_cli_args(args)
# Get all fields from the dataclass
attrs = [attr.name for attr in dataclasses.fields(cls)]
logger.info(provided_args)
# Create a dictionary of attribute values, with defaults for missing attributes
kwargs = {}
for attr in attrs:
# Handle renamed attributes or those with multiple CLI names
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
if attr == 'pipeline_config':
pipeline_config = PipelineConfig.from_kwargs(provided_args)
kwargs[attr] = pipeline_config
# Use getattr with default value from the dataclass for potentially missing attributes
else:
default_value = getattr(cls, attr, None)
if getattr(args, attr, default_value) is not None:
kwargs[attr] = getattr(args, attr, default_value)
value = getattr(args, attr, default_value)
kwargs[attr] = value # type: ignore
return cls(**kwargs)
return cls(**kwargs) # type: ignore
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
@@ -689,6 +545,9 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--tracker-project-name",
type=str,
help="Project name for tracking")
parser.add_argument("--wandb-run-name",
type=str,
help="Run name for wandb")
parser.add_argument("--seed",
type=int,
default=42,
@@ -827,4 +686,16 @@ class TrainingArgs(FastVideoArgs):
type=str,
help="Master weight type")
# VSA parameters for training with dense to sparse adaption
parser.add_argument(
"--VSA-decay-rate", # decay rate, how much sparsity you want to decay each step
type=float,
default=TrainingArgs.VSA_decay_rate,
help="VSA decay rate")
parser.add_argument(
"--VSA-decay-interval-steps", # how many steps for training with current sparsity
type=int,
default=TrainingArgs.VSA_decay_interval_steps,
help="VSA decay interval steps")
return parser
+2 -2
View File
@@ -114,7 +114,7 @@ def _info(logger: Logger,
if (main_process_only and is_main_process) or (local_main_process_only
and is_local_main_process):
logger.log(logging.INFO, msg, *args, **kwargs)
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
global _warned_local_main_process, _warned_main_process
@@ -134,7 +134,7 @@ def _info(logger: Logger,
_warned_main_process = True
if not main_process_only and not local_main_process_only:
logger.log(logging.INFO, msg, *args, **kwargs)
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
class _FastvideoLogger(Logger):
+4 -4
View File
@@ -6,7 +6,7 @@ import torch
from torch import nn
from fastvideo.v1.configs.models import DiTConfig
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
# TODO
@@ -19,7 +19,7 @@ class BaseDiT(nn.Module, ABC):
num_channels_latents: int
# always supports torch_sdpa
_supported_attention_backends: Tuple[
_Backend, ...] = DiTConfig()._supported_attention_backends
AttentionBackendEnum, ...] = DiTConfig()._supported_attention_backends
def __init_subclass__(cls) -> None:
required_class_attrs = [
@@ -65,7 +65,7 @@ class BaseDiT(nn.Module, ABC):
)
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
return self._supported_attention_backends
@@ -85,7 +85,7 @@ class CachableDiT(BaseDiT):
num_channels_latents: int
# always supports torch_sdpa
_supported_attention_backends: Tuple[
_Backend, ...] = DiTConfig()._supported_attention_backends
AttentionBackendEnum, ...] = DiTConfig()._supported_attention_backends
def __init__(self, config: DiTConfig, **kwargs) -> None:
super().__init__(config, **kwargs)
+7 -5
View File
@@ -23,7 +23,7 @@ from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
unpatchify)
from fastvideo.v1.models.dits.base import CachableDiT
from fastvideo.v1.models.utils import modulate
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
class HunyuanRMSNorm(nn.Module):
@@ -96,7 +96,8 @@ class MMDoubleStreamBlock(nn.Module):
num_attention_heads: int,
mlp_ratio: float,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
prefix: str = "",
):
super().__init__()
@@ -303,7 +304,8 @@ class MMSingleStreamBlock(nn.Module):
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
prefix: str = "",
):
super().__init__()
@@ -876,8 +878,8 @@ class IndividualTokenRefinerBlock(nn.Module):
num_heads=num_attention_heads,
head_size=hidden_size // num_attention_heads,
# TODO: remove hardcode; remove STA
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA),
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA),
)
def forward(self, x, c):
+14 -12
View File
@@ -26,7 +26,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.v1.layers.visual_embedding import TimestepEmbedder
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
class PatchEmbed2D(nn.Module):
@@ -139,16 +139,17 @@ class StepVideoRMSNorm(nn.Module):
class SelfAttention(nn.Module):
def __init__(self,
hidden_dim,
head_dim,
rope_split: Tuple[int, int, int] = (64, 32, 32),
bias: bool = False,
with_rope: bool = True,
with_qk_norm: bool = True,
attn_type: str = "torch",
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)):
def __init__(
self,
hidden_dim,
head_dim,
rope_split: Tuple[int, int, int] = (64, 32, 32),
bias: bool = False,
with_rope: bool = True,
with_qk_norm: bool = True,
attn_type: str = "torch",
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)):
super().__init__()
self.head_dim = head_dim
self.hidden_dim = hidden_dim
@@ -257,7 +258,8 @@ class CrossAttention(nn.Module):
head_dim,
bias=False,
with_qk_norm=True,
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)
) -> None:
super().__init__()
self.head_dim = head_dim
+29 -26
View File
@@ -26,7 +26,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder)
from fastvideo.v1.models.dits.base import CachableDiT
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
class WanImageEmbedding(torch.nn.Module):
@@ -125,8 +125,8 @@ class WanSelfAttention(nn.Module):
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA))
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
def forward(self, x: torch.Tensor, context: torch.Tensor,
context_lens: int):
@@ -174,7 +174,8 @@ class WanI2VCrossAttention(WanSelfAttention):
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None
) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps,
supported_attention_backends)
@@ -216,17 +217,18 @@ class WanI2VCrossAttention(WanSelfAttention):
class WanTransformerBlock(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = ""):
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
@@ -358,17 +360,18 @@ class WanTransformerBlock(nn.Module):
class WanTransformerBlock_VSA(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = ""):
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
+7 -5
View File
@@ -8,12 +8,13 @@ from torch import nn
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
ImageEncoderConfig,
TextEncoderConfig)
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
class TextEncoder(nn.Module, ABC):
_supported_attention_backends: Tuple[
_Backend, ...] = TextEncoderConfig()._supported_attention_backends
AttentionBackendEnum,
...] = TextEncoderConfig()._supported_attention_backends
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__()
@@ -34,13 +35,14 @@ class TextEncoder(nn.Module, ABC):
pass
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
return self._supported_attention_backends
class ImageEncoder(nn.Module, ABC):
_supported_attention_backends: Tuple[
_Backend, ...] = ImageEncoderConfig()._supported_attention_backends
AttentionBackendEnum,
...] = ImageEncoderConfig()._supported_attention_backends
def __init__(self, config: ImageEncoderConfig) -> None:
super().__init__()
@@ -56,5 +58,5 @@ class ImageEncoder(nn.Module, ABC):
pass
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
return self._supported_attention_backends
+3 -5
View File
@@ -81,10 +81,7 @@ def get_hf_config(
return config
def get_diffusers_config(
model: str,
fastvideo_args: Optional[dict] = None,
) -> Dict[str, Any]:
def get_diffusers_config(model: str, ) -> Dict[str, Any]:
"""Gets a configuration for the given diffusers model.
Args:
@@ -105,7 +102,8 @@ def get_diffusers_config(
# Load the config directly from the file
with open(config_file) as f:
config_dict: Dict[str, Any] = json.load(f)
if "_diffusers_version" in config_dict:
config_dict.pop("_diffusers_version")
# TODO(will): apply any overrides from inference args
return config_dict
except Exception as e:
+35 -41
View File
@@ -15,8 +15,9 @@ from safetensors.torch import load_file as safetensors_load_file
from transformers import AutoImageProcessor, AutoTokenizer
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.v1.configs.models import EncoderConfig
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
from fastvideo.v1.models.loader.fsdp_load import maybe_load_fsdp_model
@@ -45,7 +46,7 @@ class ComponentLoader(ABC):
Args:
model_path: Path to the component model
architecture: Architecture of the component model
fastvideo_args: Inference arguments
fastvideo_args: FastVideoArgs
Returns:
The loaded component
@@ -183,9 +184,10 @@ class TextEncoderLoader(ComponentLoader):
self,
model_config: Any,
model: nn.Module,
model_path: str,
) -> Generator[Tuple[str, torch.Tensor], None, None]:
primary_weights = TextEncoderLoader.Source(
model_config.model,
model_path,
prefix="",
fall_back_to_pt=getattr(model, "fall_back_to_pt_during_load", True),
allow_patterns_overrides=getattr(model, "allow_patterns_overrides",
@@ -209,8 +211,7 @@ class TextEncoderLoader(ComponentLoader):
# revision=fastvideo_args.revision,
# model_override_args=None,
# )
with open(os.path.join(model_path, "config.json")) as f:
model_config = json.load(f)
model_config = get_diffusers_config(model=model_path)
model_config.pop("_name_or_path", None)
model_config.pop("transformers_version", None)
model_config.pop("model_type", None)
@@ -220,13 +221,17 @@ class TextEncoderLoader(ComponentLoader):
# @TODO(Wei): Better way to handle this?
try:
encoder_config = fastvideo_args.text_encoder_configs[0]
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[
0]
encoder_config.update_model_arch(model_config)
encoder_precision = fastvideo_args.text_encoder_precisions[0]
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
0]
except Exception:
encoder_config = fastvideo_args.text_encoder_configs[1]
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[
1]
encoder_config.update_model_arch(model_config)
encoder_precision = fastvideo_args.text_encoder_precisions[1]
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
1]
target_device = get_torch_device()
# TODO(will): add support for other dtypes
@@ -235,7 +240,7 @@ class TextEncoderLoader(ComponentLoader):
def load_model(self,
model_path: str,
model_config,
model_config: EncoderConfig,
target_device: torch.device,
dtype: str = "fp16"):
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
@@ -245,9 +250,8 @@ class TextEncoderLoader(ComponentLoader):
model = model_cls(model_config)
weights_to_load = {name for name, _ in model.named_parameters()}
model_config.model = model_path
loaded_weights = model.load_weights(
self._get_all_weights(model_config, model))
self._get_all_weights(model_config, model, model_path))
self.counter_after_loading_weights = time.perf_counter()
logger.info(
"Loading weights took %.2f seconds",
@@ -261,7 +265,6 @@ class TextEncoderLoader(ComponentLoader):
raise ValueError("Following weights were not initialized from "
f"checkpoint: {weights_not_loaded}")
# TODO(will): add support for training/finetune
return model.eval()
@@ -284,13 +287,14 @@ class ImageEncoderLoader(TextEncoderLoader):
model_config.pop("model_type", None)
logger.info("HF Model config: %s", model_config)
encoder_config = fastvideo_args.image_encoder_config
encoder_config = fastvideo_args.pipeline_config.image_encoder_config
encoder_config.update_model_arch(model_config)
target_device = get_torch_device()
# TODO(will): add support for other dtypes
return self.load_model(model_path, encoder_config, target_device,
fastvideo_args.image_encoder_precision)
return self.load_model(
model_path, encoder_config, target_device,
fastvideo_args.pipeline_config.image_encoder_precision)
class ImageProcessorLoader(ComponentLoader):
@@ -332,18 +336,17 @@ class VAELoader(ComponentLoader):
def load(self, model_path: str, architecture: str,
fastvideo_args: FastVideoArgs):
"""Load the VAE based on the model path, architecture, and inference args."""
# TODO(will): move this to a constants file
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name")
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
config.pop("_diffusers_version")
vae_config = fastvideo_args.vae_config
vae_config = fastvideo_args.pipeline_config.vae_config
vae_config.update_model_arch(config)
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(get_torch_device())
with set_default_torch_dtype(PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]):
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(get_torch_device())
# Find all safetensors files
safetensors_list = glob.glob(
@@ -355,10 +358,8 @@ class VAELoader(ComponentLoader):
loaded = safetensors_load_file(safetensors_list[0])
vae.load_state_dict(
loaded, strict=False) # We might only load encoder or decoder
dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae = vae.eval().to(dtype)
return vae
return vae.eval()
class TransformerLoader(ComponentLoader):
@@ -374,10 +375,9 @@ class TransformerLoader(ComponentLoader):
raise ValueError(
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
config.pop("_diffusers_version")
# Config from Diffusers supersedes fastvideo's model config
dit_config = fastvideo_args.dit_config
dit_config = fastvideo_args.pipeline_config.dit_config
dit_config.update_model_arch(config)
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
@@ -391,14 +391,8 @@ class TransformerLoader(ComponentLoader):
logger.info("Loading model from %s safetensors files in %s",
len(safetensors_list), model_path)
# initialize_sequence_parallel_group(fastvideo_args.sp_size)
if fastvideo_args.training_mode:
assert isinstance(
fastvideo_args, TrainingArgs
), "fastvideo_args must be a TrainingArgs object when training_mode is True"
default_dtype = PRECISION_TO_TYPE[fastvideo_args.master_weight_type]
else:
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
default_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.dit_precision]
# Load the model using FSDP loader
logger.info("Loading model from %s, default_dtype: %s", cls_name,
@@ -462,15 +456,15 @@ class SchedulerLoader(ComponentLoader):
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 fastvideo_args.flow_shift is not None:
scheduler.set_shift(fastvideo_args.flow_shift)
if fastvideo_args.timesteps_scale is not None:
scheduler.set_timesteps_scale(fastvideo_args.timesteps_scale)
if fastvideo_args.pipeline_config.flow_shift is not None:
scheduler.set_shift(fastvideo_args.pipeline_config.flow_shift)
if fastvideo_args.pipeline_config.timesteps_scale is not None:
scheduler.set_timesteps_scale(
fastvideo_args.pipeline_config.timesteps_scale)
return scheduler
@@ -529,7 +523,7 @@ class PipelineComponentLoader:
component_model_path: Path to the component model
transformers_or_diffusers: Whether the module is from transformers or diffusers
architecture: Architecture of the component model
fastvideo_args: Inference arguments
pipeline_args: Inference arguments
Returns:
The loaded module
+1 -1
View File
@@ -49,7 +49,7 @@ def build_pipeline(fastvideo_args: FastVideoArgs) -> PipelineWithLoRA:
pipeline_architecture)
# instantiate the pipeline
pipeline = pipeline_cls(model_path, fastvideo_args, config)
pipeline = pipeline_cls(model_path, fastvideo_args)
logger.info("Pipeline instantiated")
# pipeline is now initialized and ready to use
@@ -8,13 +8,11 @@ This module defines the base class for pipelines that are composed of multiple s
import argparse
import os
from abc import ABC, abstractmethod
from copy import deepcopy
from typing import Any, Dict, List, Optional, Union, cast
import torch
from fastvideo.v1.configs.pipelines import (PipelineConfig,
get_pipeline_config_cls_for_name)
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.distributed import (
maybe_init_distributed_environment_and_model_parallel)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
@@ -22,7 +20,7 @@ from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import PipelineStage
from fastvideo.v1.utils import (maybe_download_model, shallow_asdict,
from fastvideo.v1.utils import (maybe_download_model,
verify_model_config_and_directory)
logger = init_logger(__name__)
@@ -46,24 +44,16 @@ class ComposedPipelineBase(ABC):
# TODO(will): args should support both inference args and training args
def __init__(self,
model_path: str,
fastvideo_args: FastVideoArgs,
config: Optional[Dict[str, Any]] = None,
fastvideo_args: Union[FastVideoArgs, TrainingArgs],
required_config_modules: Optional[List[str]] = None,
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None):
"""
Initialize the pipeline. After __init__, the pipeline should be ready to
use. The pipeline should be stateless and not hold any batch state.
"""
self.fastvideo_args = fastvideo_args
if fastvideo_args.training_mode:
assert isinstance(fastvideo_args, TrainingArgs)
self.training_args = fastvideo_args
assert self.training_args is not None
else:
self.fastvideo_args = fastvideo_args
assert self.fastvideo_args is not None
self.model_path = model_path
self.model_path: str = model_path
self._stages: List[PipelineStage] = []
self._stage_name_mapping: Dict[str, PipelineStage] = {}
@@ -74,13 +64,6 @@ class ComposedPipelineBase(ABC):
raise NotImplementedError(
"Subclass must set _required_config_modules")
if config is None:
# Load configuration
logger.info("Loading pipeline configuration...")
self.config = self._load_config(model_path)
else:
self.config = config
maybe_init_distributed_environment_and_model_parallel(
fastvideo_args.tp_size, fastvideo_args.sp_size)
@@ -89,6 +72,8 @@ class ComposedPipelineBase(ABC):
self.modules = self.load_modules(fastvideo_args, loaded_modules)
if fastvideo_args.training_mode:
assert isinstance(fastvideo_args, TrainingArgs)
self.training_args = fastvideo_args
assert self.training_args is not None
if self.training_args.log_validation:
self.initialize_validation_pipeline(self.training_args)
@@ -127,39 +112,16 @@ class ComposedPipelineBase(ABC):
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
"""
config = None
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, PipelineConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = kwargs
else:
config_args = shallow_asdict(config)
config_args.update(kwargs)
if args is None or args.inference_mode:
fastvideo_args = FastVideoArgs(model_path=model_path, **config_args)
fastvideo_args.model_path = model_path
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
kwargs['model_path'] = model_path
fastvideo_args = FastVideoArgs.from_kwargs(kwargs)
else:
assert args is not None, "args must be provided for training mode"
fastvideo_args = TrainingArgs.from_cli_args(args)
# TODO(will): fix this so that its not so ugly
fastvideo_args.model_path = model_path
for key, value in config_args.items():
for key, value in kwargs.items():
setattr(fastvideo_args, key, value)
fastvideo_args.use_cpu_offload = False
@@ -170,7 +132,7 @@ class ComposedPipelineBase(ABC):
# use FSDP2's MixedPrecisionPolicy to set the precision for the
# fwd, bwd, and other operations' precision.
# fastvideo_args.precision = fastvideo_args.master_weight_type
assert fastvideo_args.master_weight_type == 'fp32', 'only fp32 is supported for training'
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
# assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
@@ -250,20 +212,21 @@ class ComposedPipelineBase(ABC):
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
"""
logger.info("Loading pipeline modules from config: %s", self.config)
modules_config = deepcopy(self.config)
model_index = self._load_config(self.model_path)
logger.info("Loading pipeline modules from config: %s", model_index)
# remove keys that are not pipeline modules
modules_config.pop("_class_name")
modules_config.pop("_diffusers_version")
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
# some sanity checks
assert len(
modules_config
model_index
) > 1, "model_index.json must contain at least one pipeline module"
for module_name in self.required_config_modules:
if module_name not in modules_config:
if module_name not in model_index:
raise ValueError(
f"model_index.json must contain a {module_name} module")
@@ -273,7 +236,7 @@ class ComposedPipelineBase(ABC):
modules = {}
for module_name, (transformers_or_diffusers,
architecture) in modules_config.items():
architecture) in model_index.items():
if module_name not in required_modules:
logger.info("Skipping module %s", module_name)
continue
+7 -6
View File
@@ -36,16 +36,17 @@ class LoRAPipeline(ComposedPipelineBase):
"transformer"].config.arch_config.exclude_lora_layers
self.convert_to_lora_layers()
if self.fastvideo_args.lora_path is not None:
if self.fastvideo_args.pipeline_config.lora_path is not None:
self.set_lora_adapter(
self.fastvideo_args.lora_nickname, # type: ignore
self.fastvideo_args.lora_path)
self.fastvideo_args.pipeline_config.
lora_nickname, # type: ignore
self.fastvideo_args.pipeline_config.lora_path)
def is_target_layer(self, module_name: str) -> bool:
if self.fastvideo_args.lora_target_names is None:
if self.fastvideo_args.pipeline_config.lora_target_names is None:
return True
return any(target_name in module_name
for target_name in self.fastvideo_args.lora_target_names)
return any(target_name in module_name for target_name in
self.fastvideo_args.pipeline_config.lora_target_names)
def convert_to_lora_layers(self) -> None:
"""
@@ -67,6 +67,7 @@ class ForwardBatch:
# Latent tensors
latents: Optional[torch.Tensor] = None
raw_latent_shape: Optional[torch.Tensor] = None
noise_pred: Optional[torch.Tensor] = None
image_latent: Optional[torch.Tensor] = None
@@ -15,7 +15,7 @@ from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.pipelines.preprocess_pipeline_base import (
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
@@ -87,7 +87,7 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
"clip_feature_dtype": "",
})
return record
return record # type: ignore
EntryClass = PreprocessPipeline_I2V
@@ -6,7 +6,7 @@ This module contains an implementation of the T2V Data Preprocessing pipeline
using the modular pipeline architecture.
"""
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.v1.pipelines.preprocess_pipeline_base import (
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
@@ -1,37 +1,38 @@
import argparse
import json
import os
import torch
import torch.distributed as dist
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import maybe_download_model, shallow_asdict
from fastvideo.v1.distributed import maybe_init_distributed_environment_and_model_parallel, get_world_size
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo import PipelineConfig
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_i2v import PreprocessPipeline_I2V
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_t2v import PreprocessPipeline_T2V
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.distributed import (
get_world_size, maybe_init_distributed_environment_and_model_parallel)
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_i2v import (
PreprocessPipeline_I2V)
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.v1.utils import maybe_download_model
logger = init_logger(__name__)
def main(args):
args.model_path = maybe_download_model(args.model_path)
maybe_init_distributed_environment_and_model_parallel(args.tp_size, args.sp_size)
def main(args) -> None:
args.model_path = maybe_download_model(args.model_path)
maybe_init_distributed_environment_and_model_parallel(1, 1)
num_gpus = int(os.environ["WORLD_SIZE"])
assert num_gpus == 1, "Only support 1 GPU"
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"use_cpu_offload": False,
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
}
pipeline_config_args = shallow_asdict(pipeline_config)
pipeline_config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=args.model_path,
num_gpus=get_world_size(),
**pipeline_config_args,
)
pipeline_config.update_config_from_dict(kwargs)
fastvideo_args = FastVideoArgs(
model_path=args.model_path,
num_gpus=get_world_size(),
pipeline_config=pipeline_config,
)
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
@@ -49,7 +50,8 @@ if __name__ == "__main__":
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--preprocess_video_batch_size",
@@ -63,18 +65,15 @@ if __name__ == "__main__":
default=8,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--samples_per_file",
type=int,
default=64
)
parser.add_argument(
"--flush_frequency",
type=int,
default=256,
help="how often to save to parquet files"
)
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--samples_per_file", type=int, default=64)
parser.add_argument("--flush_frequency",
type=int,
default=256,
help="how often to save to parquet files")
parser.add_argument("--num_latent_t",
type=int,
default=28,
help="Number of latent timesteps.")
parser.add_argument("--max_height", type=int, default=480)
parser.add_argument("--max_width", type=int, default=848)
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
@@ -88,15 +87,18 @@ if __name__ == "__main__":
parser.add_argument("--speed_factor", type=float, default=1.0)
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
# text encoder & vae & diffusion model
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--text_encoder_name",
type=str,
default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
help=
"The output directory where the model predictions and checkpoints will be written.",
)
args = parser.parse_args()
main(args)
main(args)
+3 -2
View File
@@ -54,7 +54,8 @@ class DecodingStage(PipelineStage):
image = latents
else:
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (vae_dtype != torch.float32
) and not fastvideo_args.disable_autocast
@@ -77,7 +78,7 @@ class DecodingStage(PipelineStage):
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.vae_tiling:
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
+15 -12
View File
@@ -21,7 +21,7 @@ from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
st_attn_available = False
if importlib.util.find_spec("st_attn") is not None:
@@ -54,10 +54,11 @@ class DenoisingStage(PipelineStage):
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
_Backend.VIDEO_SPARSE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
supported_attention_backends=(
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
) # hack
)
def forward(
@@ -194,13 +195,15 @@ class DenoisingStage(PipelineStage):
# Prepare inputs for transformer
t_expand = t.repeat(latent_model_input.shape[0])
guidance_expand = (torch.tensor(
[fastvideo_args.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=get_torch_device(),
).to(target_dtype) * 1000.0 if fastvideo_args.embedded_cfg_scale
is not None else None)
guidance_expand = (
torch.tensor(
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=get_torch_device(),
).to(target_dtype) *
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
is not None else None)
# Predict noise residual
with torch.autocast(device_type="cuda",
+3 -2
View File
@@ -75,7 +75,8 @@ class EncodingStage(PipelineStage):
dtype=torch.float32)
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
@@ -83,7 +84,7 @@ class EncodingStage(PipelineStage):
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.vae_tiling:
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
@@ -75,10 +75,10 @@ class LatentPreparationStage(PipelineStage):
batch_size,
self.transformer.num_channels_latents,
num_frames,
height //
fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
width //
fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
height // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
width // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
)
# Validate generator if it's a list
@@ -103,6 +103,7 @@ class LatentPreparationStage(PipelineStage):
# Update batch with prepared latents
batch.latents = latents
batch.raw_latent_shape = latents.shape
return batch
@@ -119,9 +120,9 @@ class LatentPreparationStage(PipelineStage):
The batch with adjusted video length.
"""
video_length = batch.num_frames
use_temporal_scaling_frames = fastvideo_args.vae_config.use_temporal_scaling_frames
use_temporal_scaling_frames = fastvideo_args.pipeline_config.vae_config.use_temporal_scaling_frames
if use_temporal_scaling_frames:
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
temporal_scale_factor = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
latent_num_frames = (video_length - 1) // temporal_scale_factor + 1
else: # stepvideo only
latent_num_frames = video_length // 17 * 3
@@ -29,9 +29,9 @@ class StepvideoPromptEncodingStage(PipelineStage):
def forward(self, batch: ForwardBatch, fastvideo_args) -> ForwardBatch:
prompts = [batch.prompt + fastvideo_args.pos_magic]
prompts = [batch.prompt + fastvideo_args.pipeline_config.pos_magic]
bs = len(prompts)
prompts += [fastvideo_args.neg_magic] * bs
prompts += [fastvideo_args.pipeline_config.neg_magic] * bs
with set_forward_context(current_timestep=0, attn_metadata=None):
y, y_mask = self.stepllm(prompts)
clip_emb, _ = self.clip(prompts)
@@ -53,13 +53,13 @@ class TextEncodingStage(PipelineStage):
"""
assert len(self.tokenizers) == len(self.text_encoders)
assert len(self.text_encoders) == len(
fastvideo_args.text_encoder_configs)
fastvideo_args.pipeline_config.text_encoder_configs)
for tokenizer, text_encoder, encoder_config, preprocess_func, postprocess_func in zip(
self.tokenizers, self.text_encoders,
fastvideo_args.text_encoder_configs,
fastvideo_args.preprocess_text_funcs,
fastvideo_args.postprocess_text_funcs):
fastvideo_args.pipeline_config.text_encoder_configs,
fastvideo_args.pipeline_config.preprocess_text_funcs,
fastvideo_args.pipeline_config.postprocess_text_funcs):
if fastvideo_args.use_cpu_offload:
text_encoder = text_encoder.to(get_torch_device())
@@ -9,7 +9,6 @@ using the modular pipeline architecture.
"""
import os
from copy import deepcopy
from typing import Any, Dict
import torch
@@ -102,21 +101,21 @@ class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
"""
Load the modules from the config.
"""
logger.info("Loading pipeline modules from config: %s", self.config)
modules_config = deepcopy(self.config)
model_index = self._load_config(self.model_path)
logger.info("Loading pipeline modules from config: %s", model_index)
# remove keys that are not pipeline modules
modules_config.pop("_class_name")
modules_config.pop("_diffusers_version")
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
# some sanity checks
assert len(
modules_config
model_index
) > 1, "model_index.json must contain at least one pipeline module"
required_modules = ["transformer", "scheduler", "vae"]
for module_name in required_modules:
if module_name not in modules_config:
if module_name not in model_index:
raise ValueError(
f"model_index.json must contain a {module_name} module")
logger.info("Diffusers config passed sanity checks")
@@ -124,7 +123,7 @@ class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
# all the component models used by the pipeline
modules = {}
for module_name, (transformers_or_diffusers,
architecture) in modules_config.items():
architecture) in model_index.items():
component_model_path = os.path.join(self.model_path, module_name)
module = PipelineComponentLoader.load_module(
module_name=module_name,
@@ -32,7 +32,7 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
+2 -2
View File
@@ -32,7 +32,7 @@ class WanPipeline(LoRAPipeline, ComposedPipelineBase):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# We use UniPCMScheduler from Wan2.1 official repo, not the one in diffusers.
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
@@ -75,7 +75,7 @@ class WanValidationPipeline(ComposedPipelineBase):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
+1 -1
View File
@@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Optional
from fastvideo.v1.logger import init_logger
# imported by other files, do not remove
from fastvideo.v1.platforms.interface import _Backend # noqa: F401
from fastvideo.v1.platforms.interface import AttentionBackendEnum # noqa: F401
from fastvideo.v1.platforms.interface import Platform, PlatformEnum
from fastvideo.v1.utils import resolve_obj_by_qualname
+32 -15
View File
@@ -13,8 +13,9 @@ from typing_extensions import ParamSpec
import fastvideo.v1.envs as envs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.platforms.interface import (DeviceCapability, Platform,
PlatformEnum, _Backend)
from fastvideo.v1.platforms.interface import (AttentionBackendEnum,
DeviceCapability, Platform,
PlatformEnum)
from fastvideo.v1.utils import import_pynvml
logger = init_logger(__name__)
@@ -106,75 +107,85 @@ class CudaPlatformBase(Platform):
return float(torch.cuda.max_memory_allocated(device))
@classmethod
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
def get_attn_backend_cls(cls,
selected_backend: Optional[AttentionBackendEnum],
head_size: int, dtype: torch.dtype) -> str:
# TODO(will): maybe come up with a more general interface for local attention
# if distributed is False, we always try to use Flash attn
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
envs.FASTVIDEO_ATTENTION_BACKEND)
if selected_backend == _Backend.SLIDING_TILE_ATTN:
if selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
try:
from st_attn import sliding_tile_attention # noqa: F401
from fastvideo.v1.attention.backends.sliding_tile_attn import ( # noqa: F401
SlidingTileAttentionBackend)
logger.info("Using Sliding Tile Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "SLIDING_TILE_ATTN"
return "fastvideo.v1.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info(
"Sliding Tile Attention backend is not installed. Fall back to Flash Attention."
)
elif selected_backend == _Backend.SAGE_ATTN:
elif selected_backend == AttentionBackendEnum.SAGE_ATTN:
try:
from sageattention import sageattn # noqa: F401
from fastvideo.v1.attention.backends.sage_attn import ( # noqa: F401
SageAttentionBackend)
logger.info("Using Sage Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "SAGE_ATTN"
return "fastvideo.v1.attention.backends.sage_attn.SageAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info(
"Sage Attention backend is not installed. Fall back to Flash Attention."
)
elif selected_backend == _Backend.VIDEO_SPARSE_ATTN:
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
try:
from vsa import block_sparse_attn # noqa: F401
from fastvideo.v1.attention.backends.video_sparse_attn import ( # noqa: F401
VideoSparseAttentionBackend)
logger.info("Using Video Sparse Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "VIDEO_SPARSE_ATTN"
return "fastvideo.v1.attention.backends.video_sparse_attn.VideoSparseAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info(
"Video Sparse Attention backend is not installed. Fall back to Flash Attention."
)
elif selected_backend == _Backend.TORCH_SDPA:
elif selected_backend == AttentionBackendEnum.TORCH_SDPA:
logger.info("Using Torch SDPA backend.")
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
elif selected_backend == _Backend.FLASH_ATTN or selected_backend is None:
elif selected_backend == AttentionBackendEnum.FLASH_ATTN or selected_backend is None:
pass
elif selected_backend:
raise ValueError(f"Invalid attention backend for {cls.device_name}")
target_backend = _Backend.FLASH_ATTN
target_backend = AttentionBackendEnum.FLASH_ATTN
if not cls.has_device_capability(80):
logger.info(
"Cannot use FlashAttention-2 backend for Volta and Turing "
"GPUs.")
target_backend = _Backend.TORCH_SDPA
target_backend = AttentionBackendEnum.TORCH_SDPA
elif dtype not in (torch.float16, torch.bfloat16):
logger.info(
"Cannot use FlashAttention-2 backend for dtype other than "
"torch.float16 or torch.bfloat16.")
target_backend = _Backend.TORCH_SDPA
target_backend = AttentionBackendEnum.TORCH_SDPA
# FlashAttn is valid for the model, checking if the package is
# installed.
if target_backend == _Backend.FLASH_ATTN:
if target_backend == AttentionBackendEnum.FLASH_ATTN:
try:
import flash_attn # noqa: F401
@@ -187,19 +198,25 @@ class CudaPlatformBase(Platform):
logger.info(
"Cannot use FlashAttention-2 backend for head size %d.",
head_size)
target_backend = _Backend.TORCH_SDPA
target_backend = AttentionBackendEnum.TORCH_SDPA
except ImportError:
logger.info("Cannot use FlashAttention-2 backend because the "
"flash_attn package is not found. "
"Make sure that flash_attn was built and installed "
"(on by default).")
target_backend = _Backend.TORCH_SDPA
target_backend = AttentionBackendEnum.TORCH_SDPA
if target_backend == _Backend.TORCH_SDPA:
if target_backend == AttentionBackendEnum.TORCH_SDPA:
logger.info("Using Torch SDPA backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "TORCH_SDPA"
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
logger.info("Using Flash Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "FLASH_ATTN"
return "fastvideo.v1.attention.backends.flash_attn.FlashAttentionBackend"
@classmethod
+3 -2
View File
@@ -13,7 +13,7 @@ from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class _Backend(enum.Enum):
class AttentionBackendEnum(enum.Enum):
FLASH_ATTN = enum.auto()
SLIDING_TILE_ATTN = enum.auto()
TORCH_SDPA = enum.auto()
@@ -88,7 +88,8 @@ class Platform:
return self._enum == PlatformEnum.CUDA
@classmethod
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
def get_attn_backend_cls(cls,
selected_backend: Optional[AttentionBackendEnum],
head_size: int, dtype: torch.dtype) -> str:
"""Get the attention backend class of a device."""
return ""
-2
View File
@@ -1,8 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import pytest
import torch.distributed as dist
import pytest
import torch
import numpy as np
@@ -10,6 +10,7 @@ from transformers import AutoConfig
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
load_tokenizer)
# from fastvideo.v1.models.hunyuan.text_encoder import load_text_encoder, load_tokenizer
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -40,8 +41,7 @@ def test_clip_encoder():
- Produce nearly identical outputs for the same input prompts
"""
args = FastVideoArgs(model_path="openai/clip-vit-large-patch14",
text_encoder_precisions=("fp16",),
text_encoder_configs=(CLIPTextConfig(),))
pipeline_config=PipelineConfig(text_encoder_configs=(CLIPTextConfig(),), text_encoder_precisions=("fp16",)))
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
logger.info("Loading models from %s", args.model_path)
@@ -8,6 +8,7 @@ from transformers import AutoConfig
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
load_tokenizer)
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -40,8 +41,7 @@ def test_llama_encoder():
- Produce nearly identical outputs for the same input prompts
"""
args = FastVideoArgs(model_path="meta-llama/Llama-2-7b-hf",
text_encoder_precisions=("fp16",),
text_encoder_configs=(LlamaConfig(),))
pipeline_config=PipelineConfig(text_encoder_configs=(LlamaConfig(),), text_encoder_precisions=("fp16",)))
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
@@ -6,6 +6,7 @@ import pytest
import torch
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel
from fastvideo.v1.configs.pipelines import PipelineConfig
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
@@ -39,7 +40,8 @@ def test_t5_encoder():
precision).to(device).eval()
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH, text_encoder_configs=(T5Config(),), text_encoder_precisions=(precision_str,))
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH, pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),), text_encoder_precisions=(precision_str,)))
loader = TextEncoderLoader()
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
@@ -0,0 +1,177 @@
import os
from pathlib import Path
from huggingface_hub import snapshot_download
import shutil
import subprocess
import sys
from fastvideo.v1.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
# Import the training pipeline
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
NUM_NODES = "1"
MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
# preprocessing
DATA_DIR = "data"
LOCAL_RAW_DATA_DIR = Path(os.path.join(DATA_DIR, "cats"))
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
PREPROCESSING_ENTRY_FILE_PATH = "fastvideo/v1/pipelines/preprocess/v1_preprocess.py"
LOCAL_PREPROCESSED_DATA_DIR = Path(os.path.join(DATA_DIR, "cats_preprocessed_data"))
# training
NUM_GPUS_PER_NODE_TRAINING = "4"
TRAINING_ENTRY_FILE_PATH = "fastvideo/v1/training/wan_training_pipeline.py"
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "combined_parquet_dataset")
LOCAL_VALIDATION_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "validation_parquet_dataset")
LOCAL_OUTPUT_DIR = Path(os.path.join(DATA_DIR, "outputs"))
def download_data():
# create the data dir if it doesn't exist
data_dir = Path(DATA_DIR)
if data_dir.exists():
print(f"Removing existing data directory at {data_dir}")
shutil.rmtree(data_dir)
print(f"Creating data directory at {data_dir}")
os.makedirs(data_dir)
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
try:
result = snapshot_download(
repo_id="wlsaidhi/cats-overfit-merged",
local_dir=str(LOCAL_RAW_DATA_DIR),
repo_type="dataset",
resume_download=True,
token=os.environ.get("HF_TOKEN"), # In case authentication is needed
)
print(f"Download completed successfully. Files downloaded to: {result}")
# Verify the download
if not LOCAL_RAW_DATA_DIR.exists():
raise RuntimeError(f"Download appeared to succeed but {LOCAL_RAW_DATA_DIR} does not exist")
# List downloaded files
print("Downloaded files:")
for file in LOCAL_RAW_DATA_DIR.rglob("*"):
if file.is_file():
print(f" - {file.relative_to(LOCAL_RAW_DATA_DIR)}")
except Exception as e:
print(f"Error during download: {str(e)}")
raise
def run_preprocessing():
# Run torchrun command
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
PREPROCESSING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--data_merge_path", os.path.join(LOCAL_RAW_DATA_DIR, "merge_1_sample.txt"),
"--preprocess_video_batch_size", "1",
"--max_height", "480",
"--max_width", "832",
"--num_frames", "77",
"--dataloader_num_workers", "0",
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
"--train_fps", "16",
"--validation_prompt_txt", os.path.join(LOCAL_RAW_DATA_DIR, "validation_prompt_1_sample.txt"),
"--samples_per_file", "1",
"--flush_frequency", "1",
"--video_length_tolerance_range", "5",
"--dataset", "t2v",
]
process = subprocess.run(cmd, check=True)
def run_training():
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_TRAINING,
TRAINING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--inference_mode", "False",
"--pretrained_model_name_or_path", MODEL_PATH,
"--data_path", LOCAL_TRAINING_DATA_DIR,
"--validation_prompt_dir", LOCAL_VALIDATION_DATA_DIR,
"--train_batch_size", "1",
"--num_latent_t", "8",
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
"--sp_size", "4",
"--tp_size", "4",
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", "4",
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "10",
"--gradient_accumulation_steps", "1",
"--max_train_steps", "901",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--checkpointing_steps", "6000",
"--validation_steps", "100",
"--validation_sampling_steps", "50",
"--log_validation",
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--cfg", "0.0",
"--output_dir", LOCAL_OUTPUT_DIR,
"--tracker_project_name", "wan_finetune_overfit_ci",
"--num_height", "480",
"--num_width", "832",
"--num_frames", "81",
"--validation_guidance_scale", "1.0",
"--num_euler_timesteps", "50",
"--multi_phased_distill_schedule", "4000-1",
"--weight_decay", "0.01",
"--not_apply_cfg_solver",
"--dit_precision", "fp32",
"--max_grad_norm", "1.0",
]
print(f"Running training with command: {cmd}")
process = subprocess.run(cmd, check=True)
def test_e2e_overfit_single_sample():
os.environ["WANDB_MODE"] = "online"
download_data()
run_preprocessing()
run_training()
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_v0.mp4")
print(f"reference_video_file: {reference_video_file}")
final_validation_video_file = os.path.join(LOCAL_OUTPUT_DIR, "validation_step_900_inference_steps_50_video_0.mp4")
print(f"final_validation_video_file: {final_validation_video_file}")
# Ensure both files exist
assert os.path.exists(reference_video_file), f"Reference video not found at {reference_video_file}"
assert os.path.exists(final_validation_video_file), f"Validation video not found at {final_validation_video_file}"
# Compute SSIM
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
reference_video_file,
final_validation_video_file,
use_ms_ssim=True # Using MS-SSIM for better quality assessment
)
print("\n===== SSIM Results for Step 900 Validation =====")
print(f"Mean MS-SSIM: {mean_ssim:.4f}")
print(f"Min MS-SSIM: {min_ssim:.4f}")
print(f"Max MS-SSIM: {max_ssim:.4f}")
assert max_ssim > 0.5, f"Max SSIM is below 0.5: {max_ssim}"
if __name__ == "__main__":
test_e2e_overfit_single_sample()
@@ -15,7 +15,7 @@ def setup_args():
parser = argparse.ArgumentParser(description='T5 Encoder Test')
parser.add_argument('--model_path', type=str, default="google/umt5-xxl")
parser.add_argument(
'--precision',
'--dit-precision',
type=str,
default="float32",
help='Precision to use for the model (float32, float16, bfloat16)')
@@ -0,0 +1 @@
{"_timestamp":1.7496170016478686e+09,"validation_videos_8_steps":{"_type":"videos","count":5,"videos":[{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"},{"sha256":"cd72a3d513eca6b41b03b80e6fa044ce7219c35e969d2ca20b9cb48c91e585c6","size":477837,"path":"media/videos/validation_videos_8_steps_0_cd72a3d513eca6b41b03.mp4","_type":"video-file"},{"_type":"video-file","sha256":"43d47c211a69bf0be3544738e76e7d8fa58bb108c00cb21dc34eaad3c6ce7cc3","size":409419,"path":"media/videos/validation_videos_8_steps_0_43d47c211a69bf0be354.mp4"},{"_type":"video-file","sha256":"ea674ec9e200bc97563c9d87d9dc07110c3f42c0ab7277dd237ae96dd8f90a10","size":333966,"path":"media/videos/validation_videos_8_steps_0_ea674ec9e200bc97563c.mp4"},{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"}],"captions":false},"step_time":2.5065076276659966,"_wandb":{"runtime":53},"learning_rate":1e-06,"_step":5,"_runtime":53.172758961,"grad_norm":5.65625,"train_loss":0.3915919363498688,"avg_step_time":2.8116052336990833}
@@ -0,0 +1,140 @@
import os
import sys
import subprocess
from pathlib import Path
import torch.distributed.elastic.multiprocessing.errors as errors
from torch.distributed.elastic.multiprocessing.errors import record
from torch.utils.data import DataLoader
import torch
import pytest
import wandb
import json
from huggingface_hub import snapshot_download
# Import the training pipeline
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
from fastvideo.v1.training.wan_training_pipeline import main
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
wandb_name = "test_training_loss"
reference_wandb_summary_file = "fastvideo/v1/tests/training/reference_wandb_summary.json"
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "4"
def run_worker():
"""Worker function that will be run on each GPU"""
# Create and populate args
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
# Set the arguments as they are in finetune_v1_test.sh
args = parser.parse_args([
"--model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--inference_mode", "False",
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--cache_dir", "/home/.cache",
"--data_path", "data/crush-smol_parq/combined_parquet_dataset",
"--validation_prompt_dir", "data/crush-smol_parq/validation_parquet_dataset",
"--train_batch_size", "2",
"--num_latent_t", "4",
"--num_gpus", "4",
"--sp_size", "4",
"--tp_size", "4",
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", "4",
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "1",
"--gradient_accumulation_steps", "2",
"--max_train_steps", "5",
"--learning_rate", "1e-6",
"--mixed_precision", "bf16",
"--checkpointing_steps", "30",
"--validation_steps", "10",
"--validation_sampling_steps", "8",
"--log_validation",
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--cfg", "0.0",
"--output_dir", "data/wan_finetune_test",
"--tracker_project_name", "wan_finetune_ci",
"--wandb_run_name", wandb_name,
"--num_height", "480",
"--num_width", "832",
"--num_frames", "81",
"--flow_shift", "3",
"--validation_guidance_scale", "1.0",
"--num_euler_timesteps", "50",
"--multi_phased_distill_schedule", "4000-1",
"--weight_decay", "0.01",
"--not_apply_cfg_solver",
"--dit_precision", "fp32",
"--max_grad_norm", "1.0"
])
# Call the main training function
main(args)
def test_distributed_training():
"""Test the distributed training setup"""
os.environ["WANDB_MODE"] = "offline"
data_dir = Path("data/crush-smol_parq")
if not data_dir.exists():
print(f"Downloading test dataset to {data_dir}...")
snapshot_download(
repo_id="PY007/crush-smol",
local_dir=str(data_dir),
repo_type="dataset",
local_dir_use_symlinks=False
)
# Get the current file path
current_file = Path(__file__).resolve()
# Run torchrun command
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE,
str(current_file)
]
process = subprocess.run(cmd, check=True)
summary_file = "fastvideo/v1/tests/training/reference_wandb_summary.json"
reference_wandb_summary = json.load(open(reference_wandb_summary_file))
wandb_summary = json.load(open(summary_file))
fields_and_thresholds = {
'avg_step_time': 1.0,
'grad_norm': 0.1,
'step_time': 0.5,
'train_loss': 0.001
}
failures = []
for field, threshold in fields_and_thresholds.items():
ref_value = reference_wandb_summary[field]
current_value = wandb_summary[field]
diff = abs(ref_value - current_value)
print(f"INFO: {field}, diff: {diff}, threshold: {threshold}, reference: {ref_value}, current: {current_value}")
if diff > threshold:
failures.append(f"FAILED: {field} difference {diff} exceeds threshold of {threshold} (reference: {ref_value}, current: {current_value})")
if failures:
raise AssertionError("\n".join(failures))
if __name__ == "__main__":
if os.environ.get("LOCAL_RANK") is not None:
# We're being run by torchrun
run_worker()
else:
# We're being run directly
test_distributed_training()
@@ -6,6 +6,7 @@ import os
import pytest
import torch
from fastvideo.v1.configs.pipelines.base import PipelineConfig
from fastvideo.v1.distributed.parallel_state import (
get_sp_parallel_rank,
get_sp_world_size)
@@ -62,10 +63,8 @@ def test_hunyuanvideo_distributed():
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=False,
precision=precision_str)
pipeline_config=PipelineConfig(dit_config=HunyuanVideoConfig(), dit_precision=precision_str))
args.device = torch.device(f"cuda:{LOCAL_RANK}")
args.dit_config = HunyuanVideoConfig()
args.check_fastvideo_args()
loader = TransformerLoader()
model = loader.load(TRANSFORMER_PATH, "", args)
@@ -6,6 +6,7 @@ import pytest
import torch
from diffusers import WanTransformer3DModel
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -34,10 +35,8 @@ def test_wan_transformer():
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=False,
precision=precision_str)
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
args.dit_config = WanVideoConfig()
args.check_fastvideo_args()
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, "", args).to(device, dtype=precision)
+2 -2
View File
@@ -7,6 +7,7 @@ import pytest
import torch
from safetensors.torch import load_file
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.logger import init_logger
# from fastvideo.v1.models.vaes.hunyuanvae import (
# AutoencoderKLHunyuanVideo as MyHunyuanVAE)
@@ -36,9 +37,8 @@ def test_hunyuan_vae():
device = torch.device("cuda:0")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
args = FastVideoArgs(model_path=VAE_PATH, pipeline_config=PipelineConfig(vae_config=HunyuanVAEConfig(), vae_precision=precision_str))
args.device = device
args.vae_config = HunyuanVAEConfig()
loader = VAELoader()
model = loader.load(VAE_PATH, "", args)
+2 -2
View File
@@ -6,6 +6,7 @@ import pytest
import torch
from diffusers import AutoencoderKLWan
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader
@@ -29,9 +30,8 @@ def test_wan_vae():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
args = FastVideoArgs(model_path=VAE_PATH, pipeline_config=PipelineConfig(vae_config=WanVAEConfig(), vae_precision=precision_str))
args.device = device
args.vae_config = WanVAEConfig()
loader = VAELoader()
model2 = loader.load(VAE_PATH, "", args)
+107 -101
View File
@@ -98,10 +98,10 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
training_args.train_batch_size,
num_data_workers=training_args.dataloader_num_workers,
drop_last=True,
text_padding_length=training_args.text_encoder_configs[0].
arch_config.text_len, # type: ignore[attr-defined]
seed=training_args.seed,
)
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
text_len, # type: ignore[attr-defined]
seed=training_args.seed)
self.noise_scheduler = noise_scheduler
@@ -121,7 +121,9 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
if self.global_rank == 0:
project = training_args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=training_args)
wandb.init(project=project,
config=training_args,
name=training_args.wandb_run_name)
@abstractmethod
def initialize_validation_pipeline(self, training_args: TrainingArgs):
@@ -170,6 +172,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
batch_size=1,
num_data_workers=0,
drop_last=False,
drop_first_row=sampling_param.negative_prompt is not None,
cfg_rate=training_args.cfg)
if sampling_param.negative_prompt:
_, negative_prompt_embeds, negative_prompt_attention_mask, _ = validation_dataset.get_validation_negative_prompt(
@@ -177,115 +180,118 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
transformer.eval()
# Process each validation prompt
videos: List[np.ndarray] = []
captions: List[str | None] = []
for _, embeddings, masks, infos in validation_dataloader:
captions.extend([None]) # TODO(peiyuan): add caption
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
validation_steps = [step for step in validation_steps if step > 0]
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# Process each validation prompt for each validation step
for num_inference_steps in validation_steps:
step_videos: List[np.ndarray] = []
step_captions: List[str | None] = []
temporal_compression_factor = training_args.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
for _, embeddings, masks, infos in validation_dataloader:
step_captions.extend([None]) # TODO(peiyuan): add caption
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
# Prepare batch for validation
batch = ForwardBatch(
data_type="video",
latents=None,
seed=validation_seed, # Use deterministic seed
generator=torch.Generator(
device="cpu").manual_seed(validation_seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
# make sure we use the same height, width, and num_frames as the training pipeline
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
# TODO(will): validation_sampling_steps and
# validation_guidance_scale are actually passed in as a list of
# values, like "10,20,30". The validation should be run for each
# combination of values.
# num_inference_steps=fastvideo_args.validation_sampling_steps,
num_inference_steps=sampling_param.num_inference_steps,
# guidance_scale=fastvideo_args.validation_guidance_scale,
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
)
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# Run validation inference
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# Re-enable gradients for training
transformer.requires_grad_(True)
transformer.train()
# Prepare batch for validation
batch = ForwardBatch(
data_type="video",
latents=None,
seed=validation_seed, # Use deterministic seed
generator=torch.Generator(
device="cpu").manual_seed(validation_seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
num_inference_steps=
num_inference_steps, # Use the current validation step
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
if self.rank_in_sp_group != 0:
continue
# Run validation inference
with torch.no_grad(), torch.autocast("cuda",
dtype=torch.bfloat16):
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
videos.append(frames)
if self.rank_in_sp_group != 0:
continue
# Log validation results
world_group = get_world_group()
num_sp_groups = world_group.world_size // self.sp_group.world_size
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
step_videos.append(frames)
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
# results to global rank 0
if self.rank_in_sp_group == 0:
if self.global_rank == 0:
# Global rank 0 collects results from all sp_group leaders
all_videos = videos # Start with own results
all_captions = captions
# Log validation results for this step
world_group = get_world_group()
num_sp_groups = world_group.world_size // self.sp_group.world_size
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
# results to global rank 0
if self.rank_in_sp_group == 0:
if self.global_rank == 0:
# Global rank 0 collects results from all sp_group leaders
all_videos = step_videos # Start with own results
all_captions = step_captions
video_filenames = []
for i, (video,
caption) in enumerate(zip(all_videos, all_captions)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_video_{i}.mp4")
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
logs = {
"validation_videos": [
wandb.Video(filename, caption=caption) for filename,
caption in zip(video_filenames, all_captions)
]
}
wandb.log(logs, step=global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(videos, dst=0)
world_group.send_object(captions, dst=0)
video_filenames = []
for i, (video,
caption) in enumerate(zip(all_videos,
all_captions)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
)
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
logs = {
f"validation_videos_{num_inference_steps}_steps": [
wandb.Video(filename, caption=caption)
for filename, caption in zip(
video_filenames, all_captions)
]
}
wandb.log(logs, step=global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(step_videos, dst=0)
world_group.send_object(step_captions, dst=0)
# Re-enable gradients for training
transformer.train()
gc.collect()
torch.cuda.empty_cache()
+47 -5
View File
@@ -1,4 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
import importlib.util
import random
import sys
import time
@@ -10,6 +11,9 @@ import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from tqdm.auto import tqdm
import fastvideo.v1.envs as envs
from fastvideo.v1.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadata)
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
get_torch_device, get_world_group)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
@@ -27,6 +31,10 @@ from fastvideo.v1.training.training_utils import (
import wandb # isort: skip
vsa_available = False
if importlib.util.find_spec("vsa") is not None:
vsa_available = True
logger = init_logger(__name__)
# Manual gradient checking flag - set to True to enable gradient verification
@@ -41,7 +49,7 @@ class WanTrainingPipeline(TrainingPipeline):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
@@ -54,7 +62,7 @@ class WanTrainingPipeline(TrainingPipeline):
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.vae_config.load_encoder = False
args_copy.pipeline_config.vae_config.load_encoder = False
validation_pipeline = WanValidationPipeline.from_pretrained(
training_args.model_path,
args=None,
@@ -66,7 +74,7 @@ class WanTrainingPipeline(TrainingPipeline):
self.validation_pipeline = validation_pipeline
def train_one_step(
def train_one_step( # type: ignore[override]
self,
transformer,
model_type,
@@ -83,6 +91,8 @@ class WanTrainingPipeline(TrainingPipeline):
logit_mean,
logit_std,
mode_scale,
patch_size,
current_vsa_sparsity,
) -> tuple[float, float]:
assert self.training_args is not None
self.modules["transformer"].requires_grad_(True)
@@ -109,6 +119,13 @@ class WanTrainingPipeline(TrainingPipeline):
get_torch_device(), dtype=torch.bfloat16)
latents = shard_latents_across_sp(
latents, num_latent_t=self.training_args.num_latent_t)
dit_seq_shape = [
latents.shape[2] // patch_size[0],
latents.shape[3] // patch_size[1],
latents.shape[4] // patch_size[2]
]
latents = normalize_dit_input(model_type, latents)
batch_size = latents.shape[0]
noise = torch.randn_like(latents)
@@ -148,8 +165,17 @@ class WanTrainingPipeline(TrainingPipeline):
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
attn_metadata = VideoSparseAttentionMetadata(
current_timestep=timesteps,
dit_seq_shape=dit_seq_shape,
VSA_sparsity=current_vsa_sparsity)
else:
attn_metadata = None
with set_forward_context(current_timestep=timesteps,
attn_metadata=None):
attn_metadata=attn_metadata):
model_pred = transformer(**input_kwargs)
if precondition_outputs:
@@ -273,11 +299,23 @@ class WanTrainingPipeline(TrainingPipeline):
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage before train_one_step: %s MB",
gpu_memory_usage)
logger.info("VSA validation sparsity: %s",
self.training_args.VSA_sparsity)
self._log_validation(self.transformer, self.training_args, 1)
if vsa_available:
vsa_sparsity = self.training_args.VSA_sparsity
vsa_decay_rate = self.training_args.VSA_decay_rate
vsa_decay_interval_steps = self.training_args.VSA_decay_interval_steps
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
if vsa_available:
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
else:
current_vsa_sparsity = 0.0
loss, grad_norm = self.train_one_step(
self.transformer,
# args.model_type,
@@ -295,6 +333,8 @@ class WanTrainingPipeline(TrainingPipeline):
self.training_args.logit_mean,
self.training_args.logit_std,
self.training_args.mode_scale,
self.training_args.pipeline_config.dit_config.patch_size,
current_vsa_sparsity,
)
step_time = time.perf_counter() - start_time
@@ -322,6 +362,7 @@ class WanTrainingPipeline(TrainingPipeline):
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
},
step=step,
)
@@ -338,6 +379,7 @@ class WanTrainingPipeline(TrainingPipeline):
logger.info("GPU memory usage after validation: %s MB",
gpu_memory_usage)
wandb.finish()
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir,
self.training_args.max_train_steps, self.optimizer,
+2 -2
View File
@@ -238,7 +238,7 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
"serve,chat,complete",
"facebook/opt-12B",
'--port', '12323',
'--tensor-parallel-size', '4',
'--tp-size', '4',
'-tp', '2'
]
```
@@ -291,7 +291,7 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
returns:
processed_args: list[str] = [
'--port': '12323',
'--tensor-parallel-size': '4',
'--tp-size': '4',
'--vae-config.load-encoder': 'false',
'--vae-config.load-decoder': 'true'
]
+29
View File
@@ -0,0 +1,29 @@
from abc import ABC, abstractmethod
from typing import Dict
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.pipelines import ComposedPipelineBase, build_pipeline
class WorkflowBase(ABC):
pipeline_configs: Dict[str, FastVideoArgs] = {}
pipelines: Dict[str, ComposedPipelineBase] = {}
def __init__(self, fastvideo_args: FastVideoArgs):
self.fastvideo_args = fastvideo_args
def register_pipelines(self, pipeline_configs: Dict[str, FastVideoArgs]):
self.pipeline_configs.update(pipeline_configs)
def load_pipelines(self):
for pipeline_name, pipeline_config in self.pipeline_configs.items():
pipeline = build_pipeline(pipeline_config)
self.pipelines[pipeline_name] = pipeline
@abstractmethod
def get_components(self):
pass
@abstractmethod
def run(self):
pass
+112 -113
View File
@@ -1,138 +1,137 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import json
import os
import time
from multiprocessing import Pool, cpu_count
from pathlib import Path
import cv2
import torchvision
from tqdm import tqdm
def get_video_info(video_path, prompt_text):
"""Extract video information using OpenCV and corresponding prompt text"""
cap = cv2.VideoCapture(str(video_path))
def get_video_info(video_path):
"""Get video information using torchvision."""
# Read video tensor (T, C, H, W)
video_tensor, _, info = torchvision.io.read_video(str(video_path),
output_format="TCHW",
pts_unit="sec")
if not cap.isOpened():
print(f"Error: Could not open video {video_path}")
return None
num_frames = video_tensor.shape[0]
height = video_tensor.shape[2]
width = video_tensor.shape[3]
fps = info.get("video_fps", 0)
duration = num_frames / fps if fps > 0 else 0
# Get video properties
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
fps = cap.get(cv2.CAP_PROP_FPS)
frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
duration = frame_count / fps if fps > 0 else 0
cap.release()
# Extract name
_, _, videos_dir, video_name = str(video_path).split("/")
return {
"path": video_path.name,
"path": str(video_name),
"resolution": {
"width": width,
"height": height
},
"size": os.path.getsize(video_path),
"fps": fps,
"duration": duration,
"cap": [prompt_text]
"num_frames": num_frames
}
def read_prompt_file(prompt_path):
"""Read and return the content of a prompt file"""
try:
with open(prompt_path, 'r', encoding='utf-8') as f:
return f.read().strip()
except Exception as e:
print(f"Error reading prompt file {prompt_path}: {e}")
return None
def prepare_dataset_json(folder_path,
output_name="videos2caption.json",
num_workers=None) -> None:
"""Prepare dataset information from a folder containing videos and prompt.txt."""
folder_path = Path(folder_path)
# Read prompt file
prompt_file = folder_path / "prompt.txt"
if not prompt_file.exists():
raise FileNotFoundError(f"prompt.txt not found in {folder_path}")
with open(prompt_file) as f:
prompts = [line.strip() for line in f.readlines() if line.strip()]
# Read videos file
videos_file = folder_path / "videos.txt"
if not videos_file.exists():
raise FileNotFoundError(f"videos.txt not found in {folder_path}")
with open(videos_file) as f:
video_paths = [line.strip() for line in f.readlines() if line.strip()]
if len(prompts) != len(video_paths):
raise ValueError(
f"Number of prompts ({len(prompts)}) does not match number of videos ({len(video_paths)})"
)
# Prepare arguments for multiprocessing
process_args = [folder_path / video_path for video_path in video_paths]
# Determine number of workers
if num_workers is None:
num_workers = max(1, cpu_count() - 1) # Leave one CPU free
# Process videos in parallel
start_time = time.time()
with Pool(num_workers) as pool:
results = list(
tqdm(pool.imap(get_video_info, process_args),
total=len(process_args),
desc="Processing videos",
unit="video"))
# Combine results with prompts
dataset_info = []
for result, prompt in zip(results, prompts):
result["cap"] = [prompt]
dataset_info.append(result)
# Calculate total processing time
total_time = time.time() - start_time
total_videos = len(dataset_info)
avg_time_per_video = total_time / total_videos if total_videos > 0 else 0
print("\nProcessing completed:")
print(f"Total videos processed: {total_videos}")
print(f"Total time: {total_time:.2f} seconds")
print(f"Average time per video: {avg_time_per_video:.2f} seconds")
# Save to JSON file
output_file = folder_path / output_name
with open(output_file, 'w') as f:
json.dump(dataset_info, f, indent=2)
# Create merge.txt
merge_file = folder_path / "merge.txt"
with open(merge_file, 'w') as f:
f.write(f"{folder_path}/videos,{output_file}\n")
print(f"Dataset information saved to {output_file}")
print(f"Merge file created at {merge_file}")
def process_videos_and_prompts(video_dir_path, prompt_dir_path, verbose=False):
"""Process videos and their corresponding prompt files
Args:
video_dir_path (str): Path to directory containing video files
prompt_dir_path (str): Path to directory containing prompt files
verbose (bool): Whether to print verbose processing information
"""
video_dir = Path(video_dir_path)
prompt_dir = Path(prompt_dir_path)
processed_data = []
# Ensure directories exist
if not video_dir.exists() or not prompt_dir.exists():
print(f"Error: One or both directories do not exist:\nVideos: {video_dir}\nPrompts: {prompt_dir}")
return []
# Process each video file
for video_file in video_dir.glob('*.mp4'):
video_name = video_file.stem
prompt_file = prompt_dir / f"{video_name}.txt"
# Check if corresponding prompt file exists
if not prompt_file.exists():
print(f"Warning: No prompt file found for video {video_name}")
continue
# Read prompt content
prompt_text = read_prompt_file(prompt_file)
if prompt_text is None:
continue
# Process video and add to results
video_info = get_video_info(video_file, prompt_text)
if video_info:
processed_data.append(video_info)
return processed_data
def save_results(processed_data, output_path):
"""Save processed data to JSON file
Args:
processed_data (list): List of processed video information
output_path (str): Full path for output JSON file
"""
output_path = Path(output_path)
# Create parent directories if they don't exist
output_path.parent.mkdir(parents=True, exist_ok=True)
with open(output_path, 'w', encoding='utf-8') as f:
json.dump(processed_data, f, indent=2, ensure_ascii=False)
return output_path
def parse_args():
"""Parse command line arguments"""
import argparse
parser = argparse.ArgumentParser(description='Process videos and their corresponding prompt files')
parser.add_argument('--video_dir', '-v', required=True, help='Directory containing video files')
parser.add_argument('--prompt_dir', '-p', required=True, help='Directory containing prompt text files')
parser.add_argument('--output_path',
'-o',
required=True,
help='Full path for output JSON file (e.g., /path/to/output/videos2caption.json)')
parser.add_argument('--verbose', action='store_true', help='Print verbose processing information')
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description='Prepare video dataset information in JSON format')
parser.add_argument(
'--data_folder',
type=str,
required=True,
help='Path to the folder containing videos and prompt.txt')
parser.add_argument(
'--output',
type=str,
default='videos2caption.json',
help='Name of the output JSON file (default: videos2caption.json)')
parser.add_argument('--workers',
type=int,
default=32,
help='Number of worker processes (default: 16)')
return parser.parse_args()
if __name__ == "__main__":
# Parse command line arguments
args = parse_args()
# Process videos and prompts
processed_videos = process_videos_and_prompts(args.video_dir, args.prompt_dir, args.verbose)
if processed_videos:
# Save results
output_path = save_results(processed_videos, args.output_path)
print(f"\nProcessed {len(processed_videos)} videos")
print(f"Results saved to: {output_path}")
# Print example of processed data
print("\nExample of processed video info:")
print(json.dumps(processed_videos[0], indent=2))
else:
print("No videos were processed successfully")
prepare_dataset_json(args.data_folder, args.output, args.workers)
+2 -4
View File
@@ -14,7 +14,6 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_prompt_dir "$VALIDATION_DIR"\
--train_batch_size=4 \
@@ -22,7 +21,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--sp_size 4 \
--tp_size 4 \
--hsdp_replicate_dim 1 \
--hsdp_shards 4 \
--hsdp_shard_dim 4 \
--num_gpus $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 10\
@@ -43,11 +42,10 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--num_height 480 \
--num_width 832 \
--num_frames 81 \
--shift 3 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 0.01 \
--not_apply_cfg_solver \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--max_grad_norm 1.0
+62
View File
@@ -0,0 +1,62 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export WANDB_API_KEY='your_wandb_api_key'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export TRITON_CACHE_DIR=/tmp/triton_cache
DATA_DIR=~/train/
VALIDATION_DIR=~latents/test/
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
CHECKPOINT_PATH="$DATA_DIR/outputs/wan_finetune/checkpoint-5"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_training_pipeline.py \
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache" \
--data_path "$DATA_DIR" \
--validation_prompt_dir "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 16 \
--sp_size 1 \
--tp_size 1 \
--num_gpus $NUM_GPUS \
--hsdp_replicate_dim $NUM_GPUS \
--hsdp-shard-dim 1 \
--train_sp_batch_size 1 \
--dataloader_num_workers 4 \
--gradient_accumulation_steps 8 \
--max_train_steps 30000 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 6000 \
--validation_steps 100 \
--validation_sampling_steps "50" \
--log_validation \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--cfg 0.0 \
--output_dir "$DATA_DIR/outputs/wan_finetune" \
--tracker_project_name VSA_finetune \
--num_height 448 \
--num_width 832 \
--num_frames 61 \
--flow_shift 3 \
--validation_guidance_scale "5.0" \
--num_euler_timesteps 50 \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--weight_decay 0.01 \
--max_grad_norm 1.0 \
--VSA_decay_sparsity 0.9 \
--VSA_decay_rate 0.03 \
--VSA_decay_interval_steps 30 \
--VSA_val_sparsity 0.9
# --resume_from_checkpoint "$CHECKPOINT_PATH"
@@ -2,12 +2,12 @@
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="finetrainers/crush-smol/merge.txt"
OUTPUT_DIR="crush-smol_preprocess"
VALIDATION_PATH="assets/prompt.txt"
DATA_MERGE_PATH="mini_i2v_dataset/crush-smol_raw/merge.txt"
OUTPUT_DIR="mini_i2v_dataset/crush-smol_preprocessed"
VALIDATION_PATH="mini_i2v_dataset/crush-smol_raw/validation.txt"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess.py \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
@@ -19,6 +19,6 @@ torchrun --nproc_per_node=$GPU_NUM \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_prompt_txt $VALIDATION_PATH \
--samples_per_file 16 \
--flush_frequency 32 \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "i2v"
@@ -7,7 +7,7 @@ OUTPUT_DIR="data/crush-smol/latents"
VALIDATION_PATH="assets/prompt.txt"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess.py \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \