Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
97d4b984c9 | ||
|
|
2a8953d74d | ||
|
|
8801b10da7 | ||
|
|
6b413f2ec4 | ||
|
|
28b72694aa |
@@ -44,6 +44,11 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_nightly_test:
|
||||
description: "Run nightly-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
@@ -188,6 +193,24 @@ jobs:
|
||||
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
|
||||
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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__()
|
||||
@@ -211,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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -32,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
|
||||
@@ -47,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)
|
||||
@@ -188,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__()
|
||||
@@ -218,6 +225,7 @@ 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",
|
||||
@@ -280,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(
|
||||
@@ -287,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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -19,7 +19,7 @@ logger = init_logger(__name__)
|
||||
def main(args) -> None:
|
||||
args.model_path = maybe_download_model(args.model_path)
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
num_gpus = os.environ["WORLD_SIZE"]
|
||||
num_gpus = int(os.environ["WORLD_SIZE"])
|
||||
assert num_gpus == 1, "Only support 1 GPU"
|
||||
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
|
||||
kwargs = {
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 ""
|
||||
|
||||
Binary file not shown.
@@ -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()
|
||||
@@ -172,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(
|
||||
|
||||
@@ -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,9 +2,9 @@
|
||||
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/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
@@ -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 1 \
|
||||
--flush_frequency 1 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--preprocess_task "i2v"
|
||||
Reference in New Issue
Block a user