Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8cb2ae9d27 | ||
|
|
f76efe798e | ||
|
|
09f455233e | ||
|
|
aea300f690 | ||
|
|
b92219f6a6 | ||
|
|
a321b95a8a |
@@ -1,29 +1,23 @@
|
||||
<div align="center">
|
||||
<img src=assets/logos/logo.svg width="30%"/>
|
||||
</div>
|
||||
|
||||
<p align="center">
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/sv3MMKyv" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
| **[Documentation](https://hao-ai-lab.github.io/FastVideo)** | **[Quick Start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/)** | **[Weekly Dev Meeting](https://github.com/hao-ai-lab/FastVideo/discussions/982)** | 🟣💬 **[Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ)** |
|
||||
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
|
||||
## NEWS
|
||||
- ```2025/11/19```: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
|
||||
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
|
||||
<details>
|
||||
<summary>More</summary>
|
||||
- `2025/11/19`: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
|
||||
- `2025/08/04`: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
|
||||
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
### More News
|
||||
|
||||
</details>
|
||||
- `2025/06/14`: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- `2025/04/24`: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- `2025/02/18`: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
|
||||
- End-to-end post-training support for bidirectional and autoregressive models:
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
|
||||
- Data preprocessing pipeline for video, image, and text data
|
||||
@@ -44,6 +38,7 @@ FastVideo has the following features:
|
||||
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/hardware_support/) for full list of supported hardware and OS.
|
||||
|
||||
## Getting Started
|
||||
|
||||
We recommend using an environment manager such as `Conda` to create a clean environment:
|
||||
|
||||
```bash
|
||||
@@ -58,18 +53,20 @@ pip install fastvideo
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
|
||||
|
||||
## Sparse Distillation
|
||||
|
||||
For our sparse distillation techniques, please see our [distillation docs](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/) and check out our [blog](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
|
||||
See below for recipes and datasets:
|
||||
|
||||
| Model | Sparse Distillation | Dataset |
|
||||
|:-------------------------------------------------------------------------------------------: |:---------------------------------------------------------------------------------------------------------------: |:--------------------------------------------------------------------------------------------------------: |
|
||||
| [FastWan2.1-T2V-1.3B](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P) | [FastVideo Synthetic Wan2.1 480P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k) |
|
||||
| [FastWan2.1-T2V-14B-Preview](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-Diffusers) | Coming soon! | [FastVideo Synthetic Wan2.1 720P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x768x1280_250k) |
|
||||
| [FastWan2.2-TI2V-5B](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free) | [FastVideo Synthetic Wan2.2 720P](https://huggingface.co/datasets/FastVideo/Wan2.2-Syn-121x704x1280_32k) |
|
||||
| Model | Sparse Distillation | Dataset |
|
||||
| ------------------------------------------------------------------------------------- | --------------------------------------------------------------------------------------------------------------- | -------------------------------------------------------------------------------------------------------- |
|
||||
| [FastWan2.1-T2V-1.3B](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P) | [FastVideo Synthetic Wan2.1 480P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k) |
|
||||
| [FastWan2.2-TI2V-5B](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free) | [FastVideo Synthetic Wan2.2 720P](https://huggingface.co/datasets/FastVideo/Wan2.2-Syn-121x704x1280_32k) |
|
||||
|
||||
## Inference
|
||||
|
||||
### Generating Your First Video
|
||||
|
||||
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation/). Create a file called `example.py` with the following code:
|
||||
|
||||
```python
|
||||
@@ -113,7 +110,6 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview/)
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
|
||||
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
|
||||
|
||||
## Awesome work using FastVideo or our research projects
|
||||
|
||||
@@ -131,9 +127,11 @@ We welcome all contributions. Please check out our guide [here](https://hao-ai-l
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/899).
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned the design and reused code from the following projects: [Wan-Video](https://github.com/Wan-Video), [ThunderKittens](https://github.com/HazyResearch/ThunderKittens), [DMD2](https://github.com/tianweiy/DMD2), [diffusers](https://github.com/huggingface/diffusers), [xDiT](https://github.com/xdit-project/xDiT), [vLLM](https://github.com/vllm-project/vllm), [SGLang](https://github.com/sgl-project/sglang). We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
|
||||
If you find FastVideo useful, please consider citing our research work:
|
||||
|
||||
```bibtex
|
||||
|
||||
@@ -4,7 +4,7 @@ export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
@@ -14,7 +14,6 @@ NUM_GPUS=1
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_crush_smol"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "wan_ode_init_crush_smol"
|
||||
--max_train_steps 6000
|
||||
--train_batch_size 1
|
||||
@@ -34,7 +33,7 @@ parallel_args=(
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
@@ -51,20 +50,17 @@ dataset_args=(
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
--log-visualization
|
||||
--visualization-steps 100
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 6e-6
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig, WanVideoConfig
|
||||
|
||||
|
||||
@@ -76,8 +78,14 @@ class MatrixGameWanVideoArchConfig(WanVideoArchConfig):
|
||||
image_dim: int = 1280
|
||||
|
||||
|
||||
def _is_transformer_block(param_name: str, module: torch.nn.Module) -> bool:
|
||||
return bool("blocks" in param_name and param_name.split(".")[-1].isdigit())
|
||||
|
||||
|
||||
@dataclass
|
||||
class MatrixGameWanVideoConfig(WanVideoConfig):
|
||||
arch_config: MatrixGameWanVideoArchConfig = field(
|
||||
default_factory=MatrixGameWanVideoArchConfig)
|
||||
prefix: str = "Wan"
|
||||
_compile_conditions: list = field(
|
||||
default_factory=lambda: [_is_transformer_block])
|
||||
|
||||
@@ -219,6 +219,24 @@ class PipelineConfig:
|
||||
"Comma-separated list of denoising steps (e.g., '1000,757,522')",
|
||||
)
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}STA-mode",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}STA_mode",
|
||||
default=PipelineConfig.STA_mode.value,
|
||||
choices=[mode.value for mode in STA_Mode],
|
||||
help=
|
||||
"STA mode: STA_inference, STA_searching, STA_tuning, STA_tuning_cfg, None",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}skip-time-steps",
|
||||
type=int,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}skip_time_steps",
|
||||
default=PipelineConfig.skip_time_steps,
|
||||
help="Number of time steps to warmup (full attention) for STA",
|
||||
)
|
||||
|
||||
# Add VAE configuration arguments
|
||||
from fastvideo.configs.models.vaes.base import VAEConfig
|
||||
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
|
||||
@@ -299,6 +317,12 @@ class PipelineConfig:
|
||||
# 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)
|
||||
|
||||
# Convert STA_mode string to enum if necessary
|
||||
if isinstance(pipeline_config.STA_mode, str) and not isinstance(
|
||||
pipeline_config.STA_mode, STA_Mode):
|
||||
pipeline_config.STA_mode = STA_Mode(pipeline_config.STA_mode)
|
||||
|
||||
return pipeline_config
|
||||
|
||||
def check_pipeline_config(self) -> None:
|
||||
|
||||
@@ -148,62 +148,6 @@ class DataValidationStage(DatasetFilterStage):
|
||||
return batch
|
||||
|
||||
|
||||
class ResolutionFilterStage(DatasetFilterStage):
|
||||
"""Stage for filtering data items based on resolution constraints."""
|
||||
|
||||
def __init__(self,
|
||||
max_h_div_w_ratio: float = 17 / 16,
|
||||
min_h_div_w_ratio: float = 8 / 16,
|
||||
max_height: int = 1024,
|
||||
max_width: int = 1024):
|
||||
self.max_h_div_w_ratio = max_h_div_w_ratio
|
||||
self.min_h_div_w_ratio = min_h_div_w_ratio
|
||||
self.max_height = max_height
|
||||
self.max_width = max_width
|
||||
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Check if data item passes resolution filtering.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with resolution information
|
||||
|
||||
Returns:
|
||||
True if passes filter, False otherwise
|
||||
"""
|
||||
# Only apply to videos
|
||||
if not batch.is_video:
|
||||
return True
|
||||
|
||||
if batch.resolution is None:
|
||||
return False
|
||||
|
||||
height = batch.resolution.get("height", None)
|
||||
width = batch.resolution.get("width", None)
|
||||
if height is None or width is None:
|
||||
return False
|
||||
|
||||
# Check aspect ratio
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
|
||||
return self.filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""Process does nothing for resolution filtering - filtering is handled by should_keep."""
|
||||
return batch
|
||||
|
||||
def filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
|
||||
min_h_div_w_ratio: float) -> bool:
|
||||
"""Filter based on height/width ratio."""
|
||||
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
|
||||
|
||||
|
||||
class FrameSamplingStage(DatasetFilterStage):
|
||||
"""Stage for temporal frame sampling and indexing."""
|
||||
|
||||
@@ -328,11 +272,6 @@ class VideoTransformStage(DatasetStage):
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({batch.path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
batch.pixel_values = video
|
||||
return batch
|
||||
@@ -489,8 +428,6 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
tokenizer) -> None:
|
||||
"""Initialize all processing stages."""
|
||||
self.validation_stage = DataValidationStage()
|
||||
self.resolution_filter_stage = ResolutionFilterStage(
|
||||
max_height=args.max_height, max_width=args.max_width)
|
||||
self.frame_sampling_stage = FrameSamplingStage(
|
||||
num_frames=args.num_frames,
|
||||
train_fps=args.train_fps,
|
||||
@@ -543,7 +480,6 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
# Initialize counters
|
||||
filter_counts = {
|
||||
"validation_failed": 0,
|
||||
"resolution_failed": 0,
|
||||
"frame_sampling_failed": 0
|
||||
}
|
||||
sample_num_frames: list[int] = []
|
||||
@@ -579,10 +515,6 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
filter_counts["validation_failed"] += 1
|
||||
return False
|
||||
|
||||
if not self.resolution_filter_stage.should_keep(batch):
|
||||
filter_counts["resolution_failed"] += 1
|
||||
return False
|
||||
|
||||
if not self.frame_sampling_stage.should_keep(batch):
|
||||
filter_counts["frame_sampling_failed"] += 1
|
||||
return False
|
||||
@@ -594,10 +526,9 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
after_count: int):
|
||||
"""Log filtering statistics."""
|
||||
logger.info(
|
||||
"validation_failed: %d, resolution_failed: %d, frame_sampling_failed: %d, "
|
||||
"validation_failed: %d, frame_sampling_failed: %d, "
|
||||
"Counter(sample_num_frames): %s, before filter: %d, after filter: %d",
|
||||
filter_counts['validation_failed'],
|
||||
filter_counts['resolution_failed'],
|
||||
filter_counts['frame_sampling_failed'], Counter(sample_num_frames),
|
||||
before_count, after_count)
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ diffusion models.
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
@@ -31,6 +32,21 @@ from fastvideo.worker.executor import Executor
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _infer_latent_batch_size(batch: ForwardBatch) -> int:
|
||||
if isinstance(batch.prompt, list):
|
||||
latent_batch_size = len(batch.prompt)
|
||||
elif batch.prompt is not None:
|
||||
latent_batch_size = 1
|
||||
elif batch.prompt_embeds is not None and len(batch.prompt_embeds) > 0:
|
||||
latent_batch_size = batch.prompt_embeds[0].shape[0]
|
||||
else:
|
||||
raise ValueError(
|
||||
"Cannot infer batch size from batch; no prompt or prompt_embeds found"
|
||||
)
|
||||
latent_batch_size *= batch.num_videos_per_prompt
|
||||
return latent_batch_size
|
||||
|
||||
|
||||
class VideoGenerator:
|
||||
"""
|
||||
A unified class for generating videos using diffusion models.
|
||||
@@ -372,8 +388,31 @@ class VideoGenerator:
|
||||
|
||||
# Run inference
|
||||
start_time = time.perf_counter()
|
||||
output_batch = self.executor.execute_forward(batch, fastvideo_args)
|
||||
samples = output_batch.output
|
||||
|
||||
# Execute forward pass in a new thread for non-blocking tensor allocation
|
||||
result_container = {}
|
||||
|
||||
def execute_forward_thread():
|
||||
result_container['output_batch'] = self.executor.execute_forward(
|
||||
batch, fastvideo_args)
|
||||
|
||||
thread = threading.Thread(target=execute_forward_thread)
|
||||
thread.start()
|
||||
latent_batch_size = _infer_latent_batch_size(batch)
|
||||
samples = torch.empty((latent_batch_size, 3, sampling_param.num_frames,
|
||||
sampling_param.height, sampling_param.width),
|
||||
device='cpu',
|
||||
pin_memory=fastvideo_args.pin_cpu_memory)
|
||||
thread.join()
|
||||
|
||||
output_batch = result_container['output_batch']
|
||||
if output_batch.output.shape == samples.shape:
|
||||
samples.copy_(output_batch.output)
|
||||
else:
|
||||
logger.warning(
|
||||
"Output shape %s does not match expected shape %s; use slow path",
|
||||
output_batch.output.shape, samples.shape)
|
||||
samples = output_batch.output.cpu()
|
||||
logging_info = output_batch.logging_info
|
||||
|
||||
gen_time = time.perf_counter() - start_time
|
||||
|
||||
@@ -10,7 +10,7 @@ from enum import Enum
|
||||
from typing import Any, TYPE_CHECKING
|
||||
|
||||
from fastvideo.configs.configs import PreprocessConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, STA_Mode
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.utils import clean_cli_args
|
||||
from fastvideo.layers.quantization import QUANTIZATION_METHODS, QuantizationMethods
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -141,8 +141,6 @@ class FastVideoArgs:
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
mask_strategy_file_path: str | None = None
|
||||
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
|
||||
skip_time_steps: int = 15
|
||||
|
||||
# Compilation
|
||||
enable_torch_compile: bool = False
|
||||
@@ -465,20 +463,6 @@ class FastVideoArgs:
|
||||
)
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
parser.add_argument(
|
||||
"--STA-mode",
|
||||
type=str,
|
||||
default=FastVideoArgs.STA_mode.value,
|
||||
choices=[mode.value for mode in STA_Mode],
|
||||
help=
|
||||
"STA mode contains STA_inference, STA_searching, STA_tuning, STA_tuning_cfg, None",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip-time-steps",
|
||||
type=int,
|
||||
default=FastVideoArgs.skip_time_steps,
|
||||
help="Number of time steps to warmup (full attention) for STA",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mask-strategy-file-path",
|
||||
type=str,
|
||||
@@ -932,6 +916,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
training_state_checkpointing_steps: int = 0 # for resuming training
|
||||
weight_only_checkpointing_steps: int = 0 # for inference
|
||||
log_visualization: bool = False
|
||||
visualization_steps: int = 0
|
||||
# simulate generator forward to match inference
|
||||
simulate_generator_forward: bool = False
|
||||
warp_denoising_step: bool = False
|
||||
@@ -1095,6 +1080,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--log-validation",
|
||||
action=StoreBoolean,
|
||||
help="Whether to log validation results")
|
||||
parser.add_argument("--visualization-steps",
|
||||
type=int,
|
||||
help="Number of visualization steps")
|
||||
parser.add_argument("--tracker-project-name",
|
||||
type=str,
|
||||
help="Project name for tracking")
|
||||
|
||||
+17
-14
@@ -102,6 +102,7 @@ def _info(logger: Logger,
|
||||
- When both are False, the message will be logged from all processes
|
||||
- By default, only logs from processes with LOCAL_RANK=0
|
||||
"""
|
||||
is_distributed = int(os.environ.get("WORLD_SIZE", 1)) > 1
|
||||
try:
|
||||
local_rank = int(os.environ["LOCAL_RANK"])
|
||||
rank = int(os.environ["RANK"])
|
||||
@@ -118,20 +119,22 @@ def _info(logger: Logger,
|
||||
|
||||
global _warned_local_main_process, _warned_main_process
|
||||
|
||||
if not _warned_local_main_process and local_main_process_only:
|
||||
logger.warning(
|
||||
'%s By default, logger.info(..) will only log from the local main process. Set logger.info(..., is_local_main_process=False) to log from all processes.%s',
|
||||
GREEN,
|
||||
RESET,
|
||||
)
|
||||
_warned_local_main_process = True
|
||||
if not _warned_main_process and main_process_only:
|
||||
logger.warning(
|
||||
'%s is_main_process_only is set to True, logging only from the main (RANK==0) process.%s',
|
||||
GREEN,
|
||||
RESET,
|
||||
)
|
||||
_warned_main_process = True
|
||||
# Only show process-awareness warnings when actually running distributed
|
||||
if is_distributed:
|
||||
if not _warned_local_main_process and local_main_process_only:
|
||||
logger.warning(
|
||||
'%s By default, logger.info(..) will only log from the local main process. Set logger.info(..., is_local_main_process=False) to log from all processes.%s',
|
||||
GREEN,
|
||||
RESET,
|
||||
)
|
||||
_warned_local_main_process = True
|
||||
if not _warned_main_process and main_process_only:
|
||||
logger.warning(
|
||||
'%s is_main_process_only is set to True, logging only from the main (RANK==0) process.%s',
|
||||
GREEN,
|
||||
RESET,
|
||||
)
|
||||
_warned_main_process = True
|
||||
|
||||
if not main_process_only and not local_main_process_only:
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
|
||||
|
||||
@@ -33,7 +33,7 @@ from fastvideo.layers.visual_embedding import (PatchEmbed)
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
class CausalWanSelfAttention(nn.Module):
|
||||
@@ -286,6 +286,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -452,7 +454,6 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
This function will be run for num_frame times.
|
||||
Process the latent frames one by one (1560 tokens each)
|
||||
"""
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
|
||||
@@ -277,7 +277,7 @@ def load_video(
|
||||
if convert_method is not None:
|
||||
pil_images = convert_method(pil_images)
|
||||
|
||||
return pil_images, original_fps if return_fps else pil_images
|
||||
return (pil_images, original_fps) if return_fps else pil_images
|
||||
|
||||
|
||||
def get_default_height_width(
|
||||
|
||||
@@ -101,6 +101,58 @@ class ComposedPipelineBase(ABC):
|
||||
module.requires_grad_(True)
|
||||
module.train()
|
||||
|
||||
@staticmethod
|
||||
def _compile_with_conditions(
|
||||
module: torch.nn.Module,
|
||||
compile_kwargs: dict[str, Any],
|
||||
) -> int:
|
||||
"""Compile submodules that match module._compile_conditions."""
|
||||
compile_conditions = getattr(module, "_compile_conditions", None)
|
||||
if not compile_conditions:
|
||||
return 0
|
||||
|
||||
compiled_count = 0
|
||||
for name, submodule in module.named_modules():
|
||||
if not name:
|
||||
continue
|
||||
if any(cond(name, submodule) for cond in compile_conditions):
|
||||
submodule.forward = torch.compile(submodule.forward,
|
||||
**compile_kwargs)
|
||||
compiled_count += 1
|
||||
return compiled_count
|
||||
|
||||
def _maybe_compile_pipeline_module(
|
||||
self,
|
||||
module_name: str,
|
||||
fsdp_module_cls: type | None,
|
||||
compile_kwargs: dict[str, Any],
|
||||
) -> None:
|
||||
if module_name not in self.modules:
|
||||
return
|
||||
|
||||
module = self.modules[module_name]
|
||||
if fsdp_module_cls is not None and isinstance(module, fsdp_module_cls):
|
||||
logger.info(
|
||||
"%s is already FSDP-wrapped; skipping torch.compile in pipeline",
|
||||
module_name.capitalize(),
|
||||
)
|
||||
return
|
||||
|
||||
compiled_count = self._compile_with_conditions(module, compile_kwargs)
|
||||
if compiled_count > 0:
|
||||
logger.info(
|
||||
"Enabled torch.compile for %d submodules in %s via _compile_conditions with kwargs=%s",
|
||||
compiled_count,
|
||||
module_name,
|
||||
compile_kwargs,
|
||||
)
|
||||
return
|
||||
|
||||
# Backward-compatible fallback: compile full module if no condition matched.
|
||||
logger.info("Enabling torch.compile for %s with kwargs=%s", module_name,
|
||||
compile_kwargs)
|
||||
self.modules[module_name] = torch.compile(module, **compile_kwargs)
|
||||
|
||||
def post_init(self) -> None:
|
||||
assert self.fastvideo_args is not None, "fastvideo_args must be set"
|
||||
if self.post_init_called:
|
||||
@@ -116,7 +168,6 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
self.initialize_pipeline(self.fastvideo_args)
|
||||
if self.fastvideo_args.enable_torch_compile:
|
||||
transformer_module = self.modules["transformer"]
|
||||
if self.fastvideo_args.training_mode:
|
||||
logger.info(
|
||||
"Torch Compile enabled via FSDP loader for training; skipping additional pipeline compile"
|
||||
@@ -130,29 +181,16 @@ class ComposedPipelineBase(ABC):
|
||||
fsdp_module_cls = None
|
||||
|
||||
compile_kwargs = self.fastvideo_args.torch_compile_kwargs or {}
|
||||
if fsdp_module_cls is not None and isinstance(
|
||||
transformer_module, fsdp_module_cls):
|
||||
logger.info(
|
||||
"Transformer is already FSDP-wrapped; skipping torch.compile in pipeline"
|
||||
)
|
||||
else:
|
||||
logger.info("Enabling torch.compile for DiT with kwargs=%s",
|
||||
compile_kwargs)
|
||||
self.modules["transformer"] = torch.compile(
|
||||
transformer_module, **compile_kwargs)
|
||||
if "transformer_2" in self.modules:
|
||||
transformer_module_2 = self.modules["transformer_2"]
|
||||
if fsdp_module_cls is not None and isinstance(
|
||||
transformer_module_2, fsdp_module_cls):
|
||||
logger.info(
|
||||
"Transformer_2 is already FSDP-wrapped; skipping torch.compile in pipeline"
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
"Enabling torch.compile for Transformer_2 with kwargs=%s",
|
||||
compile_kwargs)
|
||||
self.modules["transformer_2"] = torch.compile(
|
||||
transformer_module_2, **compile_kwargs)
|
||||
self._maybe_compile_pipeline_module(
|
||||
module_name="transformer",
|
||||
fsdp_module_cls=fsdp_module_cls,
|
||||
compile_kwargs=compile_kwargs,
|
||||
)
|
||||
self._maybe_compile_pipeline_module(
|
||||
module_name="transformer_2",
|
||||
fsdp_module_cls=fsdp_module_cls,
|
||||
compile_kwargs=compile_kwargs,
|
||||
)
|
||||
logger.info("Torch Compile enabled for DiT")
|
||||
|
||||
if not self.fastvideo_args.training_mode:
|
||||
|
||||
@@ -242,8 +242,8 @@ class DecodingStage(PipelineStage):
|
||||
decoded_frames = self.decode(cur_latent, fastvideo_args)
|
||||
batch.trajectory_decoded.append(decoded_frames.cpu().float())
|
||||
|
||||
# Convert to CPU float32 for compatibility
|
||||
frames = frames.cpu().float()
|
||||
# Convert to float32 for compatibility
|
||||
frames = frames.to(torch.float32)
|
||||
|
||||
# Crop padding if this is a LongCat refinement
|
||||
if hasattr(batch, 'num_cond_frames_added') and hasattr(
|
||||
|
||||
@@ -314,6 +314,7 @@ class DenoisingStage(PipelineStage):
|
||||
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
|
||||
else:
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
t_expand = t_expand.to(get_local_torch_device())
|
||||
|
||||
use_meanflow = getattr(self.transformer.config, "use_meanflow",
|
||||
False)
|
||||
@@ -509,7 +510,7 @@ class DenoisingStage(PipelineStage):
|
||||
mgr2.release_all()
|
||||
|
||||
# Save STA mask search results if needed
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING:
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.pipeline_config.STA_mode == STA_Mode.STA_SEARCHING:
|
||||
self.save_sta_search_results(batch)
|
||||
|
||||
# deallocate transformer if on mps
|
||||
@@ -602,8 +603,8 @@ class DenoisingStage(PipelineStage):
|
||||
"""
|
||||
# TODO(kevin): STA mask search, currently only support Wan2.1 with 69x768x1280
|
||||
from fastvideo.attention.backends.STA_configuration import configure_sta
|
||||
STA_mode = fastvideo_args.STA_mode
|
||||
skip_time_steps = fastvideo_args.skip_time_steps
|
||||
STA_mode = fastvideo_args.pipeline_config.STA_mode
|
||||
skip_time_steps = fastvideo_args.pipeline_config.skip_time_steps
|
||||
if batch.timesteps is None:
|
||||
raise ValueError("Timesteps must be provided")
|
||||
timesteps_num = batch.timesteps.shape[0]
|
||||
|
||||
@@ -343,7 +343,7 @@ class SRDenoisingStage(PipelineStage):
|
||||
batch.latents = latents
|
||||
|
||||
# Save STA mask search results if needed
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING:
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.pipeline_config.STA_mode == STA_Mode.STA_SEARCHING:
|
||||
self.save_sta_search_results(batch)
|
||||
|
||||
# deallocate transformer if on mps
|
||||
|
||||
@@ -43,7 +43,10 @@ from fastvideo.training.training_utils import (
|
||||
from fastvideo.utils import (is_vsa_available, maybe_download_model,
|
||||
set_random_seed, verify_model_config_and_directory)
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -1267,7 +1270,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
with torch.no_grad():
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, training_args)
|
||||
samples = output_batch.output
|
||||
samples = output_batch.output.cpu()
|
||||
if self.rank_in_sp_group != 0:
|
||||
continue
|
||||
|
||||
|
||||
@@ -18,7 +18,10 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.utils import is_vsa_available, shallow_asdict
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
|
||||
SelfForcingFlowMatchScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
|
||||
WanCausalDMDPipeline)
|
||||
from fastvideo.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.training.training_utils import (
|
||||
@@ -57,15 +58,17 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
self.timestep_shift = self.training_args.pipeline_config.flow_shift
|
||||
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
|
||||
self.noise_scheduler = SelfForcingFlowMatchScheduler(
|
||||
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
|
||||
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
|
||||
training=True)
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
logger.info("dmd_denoising_steps: %s",
|
||||
self.training_args.pipeline_config.dmd_denoising_steps)
|
||||
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
|
||||
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250, 0],
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
|
||||
@@ -161,27 +164,12 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
device, dtype=torch.bfloat16)
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch, trajectory_latents.to(
|
||||
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
|
||||
|
||||
## TEMP Used for loading the sf .pt files directly
|
||||
"""
|
||||
self.manual_idx = self.manual_idx % 155
|
||||
path = f"/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_pt_vidprom_1000/{self.manual_idx:05d}.pt"
|
||||
logger.info("path: %s", path)
|
||||
self.manual_idx += 1
|
||||
# path = "/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_single_full/00000.pt"
|
||||
b = torch.load(path)
|
||||
training_batch.encoder_hidden_states = b["text_embedding"][0].unsqueeze(
|
||||
0).to(device, dtype=torch.bfloat16)
|
||||
trajectory_latents = b["ode_latent"].to(device, dtype=torch.bfloat16)
|
||||
logger.info("trajectory_latents: %s", trajectory_latents.shape)
|
||||
logger.info("encoder_hidden_states: %s",
|
||||
training_batch.encoder_hidden_states.shape)
|
||||
assert trajectory_latents.shape[1] <= 10, "trajectory_latents.shape[1] must be <= 10"
|
||||
return training_batch, trajectory_latents.to(
|
||||
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
|
||||
"""
|
||||
return training_batch, trajectory_latents[:, :, :self.training_args.
|
||||
num_latent_t].to(
|
||||
device,
|
||||
dtype=torch.bfloat16
|
||||
), trajectory_timesteps.to(
|
||||
device)
|
||||
|
||||
def _get_timestep(self,
|
||||
min_timestep: int,
|
||||
@@ -225,7 +213,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
|
||||
if self._cached_closest_idx_per_dmd is None:
|
||||
self._cached_closest_idx_per_dmd = torch.tensor(
|
||||
[0, 12, 24, 36], dtype=torch.long).cpu()
|
||||
[0, 12, 24, 36, S - 1], dtype=torch.long).cpu()
|
||||
# [0, 1, 2, 3], dtype=torch.long).cpu()
|
||||
logger.info("self._cached_closest_idx_per_dmd: %s",
|
||||
self._cached_closest_idx_per_dmd)
|
||||
@@ -367,8 +355,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
assert latent_key in latents_vis_dict and latents_vis_dict[
|
||||
latent_key] is not None
|
||||
latent = latents_vis_dict[latent_key]
|
||||
pixel_latent = self.validation_pipeline.decoding_stage.decode(
|
||||
latent, training_args)
|
||||
pixel_latent = self.decoding_stage.decode(latent, training_args)
|
||||
|
||||
video = pixel_latent.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
|
||||
@@ -30,7 +30,10 @@ from fastvideo.profiler import profile_region
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
|
||||
class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
|
||||
@@ -13,16 +13,18 @@ import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torchvision
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from einops import rearrange
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionMetadataBuilder)
|
||||
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
|
||||
try:
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionMetadataBuilder)
|
||||
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
|
||||
except Exception:
|
||||
pass
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset import build_parquet_map_style_dataloader
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
@@ -48,8 +50,12 @@ from fastvideo.training.training_utils import (
|
||||
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
|
||||
set_random_seed, shallow_asdict)
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
vmoba_available = is_vmoba_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
vmoba_available = is_vmoba_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
vmoba_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -108,7 +114,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
# Set random seeds for deterministic training
|
||||
assert self.seed is not None, "seed must be set"
|
||||
set_random_seed(self.seed)
|
||||
set_random_seed(self.seed + self.global_rank)
|
||||
self.transformer.train()
|
||||
if training_args.enable_gradient_checkpointing_type is not None:
|
||||
self.transformer = apply_activation_checkpointing(
|
||||
@@ -588,15 +594,15 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
round(num_trainable_params / 1e9, 3))
|
||||
|
||||
# Set random seeds for deterministic training
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
self.noise_random_generator = torch.Generator(
|
||||
device="cpu").manual_seed(self.seed + self.global_rank)
|
||||
self.noise_gen_cuda = torch.Generator(
|
||||
device=current_platform.device_name).manual_seed(self.seed)
|
||||
device=current_platform.device_name).manual_seed(self.seed +
|
||||
self.global_rank)
|
||||
self.validation_random_generator = torch.Generator(
|
||||
device="cpu").manual_seed(self.seed)
|
||||
logger.info("Initialized random seeds with seed: %s", self.seed)
|
||||
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
device="cpu").manual_seed(self.seed + self.global_rank)
|
||||
logger.info("Initialized random seeds with seed: %s",
|
||||
self.seed + self.global_rank)
|
||||
|
||||
if self.training_args.resume_from_checkpoint:
|
||||
self._resume_from_checkpoint()
|
||||
@@ -661,26 +667,31 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
"grad_norm": grad_norm,
|
||||
"vsa_sparsity": current_vsa_sparsity,
|
||||
}
|
||||
metrics["batch_size"] = int(training_batch.raw_latent_shape[0])
|
||||
try:
|
||||
metrics["batch_size"] = int(
|
||||
training_batch.raw_latent_shape[0])
|
||||
|
||||
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
|
||||
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
|
||||
training_batch.raw_latent_shape[3] //
|
||||
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
|
||||
if training_batch.encoder_hidden_states is not None:
|
||||
context_len = int(
|
||||
training_batch.encoder_hidden_states.shape[1])
|
||||
else:
|
||||
context_len = 0
|
||||
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
|
||||
seq_len = (
|
||||
training_batch.raw_latent_shape[2] // patch_t) * (
|
||||
training_batch.raw_latent_shape[3] // patch_h) * (
|
||||
training_batch.raw_latent_shape[4] // patch_w)
|
||||
if training_batch.encoder_hidden_states is not None:
|
||||
context_len = int(
|
||||
training_batch.encoder_hidden_states.shape[1])
|
||||
else:
|
||||
context_len = 0
|
||||
|
||||
metrics["dit_seq_len"] = int(seq_len)
|
||||
metrics["context_len"] = context_len
|
||||
metrics["dit_seq_len"] = int(seq_len)
|
||||
metrics["context_len"] = context_len
|
||||
|
||||
arch_config = self.training_args.pipeline_config.dit_config.arch_config
|
||||
arch_config = self.training_args.pipeline_config.dit_config.arch_config
|
||||
|
||||
metrics["hidden_dim"] = arch_config.hidden_size
|
||||
metrics["num_layers"] = arch_config.num_layers
|
||||
metrics["ffn_dim"] = arch_config.ffn_dim
|
||||
metrics["hidden_dim"] = arch_config.hidden_size
|
||||
metrics["num_layers"] = arch_config.num_layers
|
||||
metrics["ffn_dim"] = arch_config.ffn_dim
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
self.tracker.log(metrics, step)
|
||||
if step % self.training_args.training_state_checkpointing_steps == 0:
|
||||
@@ -693,12 +704,14 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
self.noise_random_generator)
|
||||
self.transformer.train()
|
||||
self.sp_group.barrier()
|
||||
|
||||
if self.training_args.log_visualization and step % self.training_args.visualization_steps == 0:
|
||||
self.visualize_intermediate_latents(training_batch,
|
||||
self.training_args, step)
|
||||
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
with self.profiler_controller.region(
|
||||
"profiler_region_training_validation"):
|
||||
if self.training_args.log_visualization:
|
||||
self.visualize_intermediate_latents(
|
||||
training_batch, self.training_args, step)
|
||||
self._log_validation(self.transformer, self.training_args,
|
||||
step)
|
||||
gpu_memory_usage = current_platform.get_torch_device(
|
||||
@@ -857,7 +870,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# Run validation inference
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, training_args)
|
||||
samples = output_batch.output
|
||||
samples = output_batch.output.cpu()
|
||||
|
||||
if self.rank_in_sp_group != 0:
|
||||
continue
|
||||
|
||||
@@ -257,8 +257,6 @@ def save_distillation_checkpoint(
|
||||
if generator_scheduler is not None:
|
||||
generator_states["scheduler"] = SchedulerWrapper(
|
||||
generator_scheduler)
|
||||
if generator_ema is not None:
|
||||
generator_states["ema"] = generator_ema.state_dict()
|
||||
|
||||
generator_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
|
||||
"generator")
|
||||
@@ -290,8 +288,6 @@ def save_distillation_checkpoint(
|
||||
if generator_scheduler_2 is not None:
|
||||
generator_2_states["scheduler"] = SchedulerWrapper(
|
||||
generator_scheduler_2)
|
||||
if generator_ema_2 is not None:
|
||||
generator_2_states["ema"] = generator_ema_2.state_dict()
|
||||
|
||||
generator_2_dcp_dir = os.path.join(save_dir,
|
||||
"distributed_checkpoint",
|
||||
@@ -417,6 +413,67 @@ def save_distillation_checkpoint(
|
||||
rank,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Persist EMA separately to avoid shape mismatches across ranks.
|
||||
# Supports:
|
||||
# - mode="rank0_full": save consolidated EMA only on rank 0
|
||||
# - mode="local_shard": save per-rank EMA shard for each rank
|
||||
try:
|
||||
if generator_ema is not None and getattr(generator_ema, "mode",
|
||||
None) == "rank0_full":
|
||||
_save_rank0_full_ema_safetensors(generator_ema,
|
||||
generator_transformer, rank,
|
||||
save_dir, "generator_ema")
|
||||
elif generator_ema is not None and getattr(generator_ema, "mode",
|
||||
None) == "local_shard":
|
||||
# Save per-rank shard
|
||||
ema_dir_shard = os.path.join(save_dir, "ema_local_shard")
|
||||
os.makedirs(ema_dir_shard, exist_ok=True)
|
||||
ema_shard_path = os.path.join(ema_dir_shard,
|
||||
f"generator_ema_rank{rank}.pt")
|
||||
torch.save(generator_ema.state_dict(), ema_shard_path)
|
||||
logger.info(
|
||||
"rank: %s, saved generator EMA shard (local_shard) to %s",
|
||||
rank,
|
||||
ema_shard_path,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Also consolidate EMA to a single full-state file on rank 0 by applying EMA to the model and gathering
|
||||
_consolidate_local_shard_ema_and_save_safetensors(
|
||||
generator_ema, generator_transformer, rank, save_dir,
|
||||
"generator_ema")
|
||||
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed saving EMA separately: %s", rank,
|
||||
str(e))
|
||||
|
||||
try:
|
||||
if generator_ema_2 is not None and getattr(generator_ema_2, "mode",
|
||||
None) == "rank0_full":
|
||||
_save_rank0_full_ema_safetensors(generator_ema_2,
|
||||
generator_transformer_2, rank,
|
||||
save_dir, "generator_ema_2")
|
||||
elif generator_ema_2 is not None and getattr(generator_ema_2, "mode",
|
||||
None) == "local_shard":
|
||||
# Save per-rank shard for EMA_2
|
||||
ema_dir_shard_2 = os.path.join(save_dir, "ema_local_shard")
|
||||
os.makedirs(ema_dir_shard_2, exist_ok=True)
|
||||
ema2_shard_path = os.path.join(ema_dir_shard_2,
|
||||
f"generator_ema_2_rank{rank}.pt")
|
||||
torch.save(generator_ema_2.state_dict(), ema2_shard_path)
|
||||
logger.info(
|
||||
"rank: %s, saved generator_2 EMA shard (local_shard) to %s",
|
||||
rank,
|
||||
ema2_shard_path,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Also consolidate EMA_2 to a single full-state file on rank 0
|
||||
_consolidate_local_shard_ema_and_save_safetensors(
|
||||
generator_ema_2, generator_transformer_2, rank, save_dir,
|
||||
"generator_ema_2")
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed saving EMA_2 separately: %s", rank,
|
||||
str(e))
|
||||
|
||||
# Save generator model weights (consolidated) for inference
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(generator_transformer,
|
||||
device=None)
|
||||
@@ -454,46 +511,45 @@ def save_distillation_checkpoint(
|
||||
logger.info("--> distillation checkpoint saved at step %s to %s", step,
|
||||
weight_path)
|
||||
|
||||
# Save generator_2 model weights (consolidated) for inference (MoE support)
|
||||
if generator_transformer_2 is not None:
|
||||
inference_save_dir_2 = os.path.join(
|
||||
save_dir, "generator_2_inference_transformer")
|
||||
cpu_state_2 = gather_state_dict_on_cpu_rank0(
|
||||
generator_transformer_2, device=None)
|
||||
# Save generator_2 model weights (consolidated) for inference (MoE support)
|
||||
if generator_transformer_2 is not None:
|
||||
inference_save_dir_2 = os.path.join(
|
||||
save_dir, "generator_2_inference_transformer")
|
||||
cpu_state_2 = gather_state_dict_on_cpu_rank0(generator_transformer_2,
|
||||
device=None)
|
||||
|
||||
if rank == 0:
|
||||
os.makedirs(inference_save_dir_2, exist_ok=True)
|
||||
weight_path_2 = os.path.join(
|
||||
inference_save_dir_2, "diffusion_pytorch_model.safetensors")
|
||||
logger.info(
|
||||
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
|
||||
rank,
|
||||
weight_path_2,
|
||||
local_main_process_only=False)
|
||||
if rank == 0:
|
||||
os.makedirs(inference_save_dir_2, exist_ok=True)
|
||||
weight_path_2 = os.path.join(inference_save_dir_2,
|
||||
"diffusion_pytorch_model.safetensors")
|
||||
logger.info(
|
||||
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
|
||||
rank,
|
||||
weight_path_2,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Convert training format to diffusers format and save
|
||||
diffusers_state_dict_2 = custom_to_hf_state_dict(
|
||||
cpu_state_2,
|
||||
generator_transformer_2.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict_2, weight_path_2)
|
||||
# Convert training format to diffusers format and save
|
||||
diffusers_state_dict_2 = custom_to_hf_state_dict(
|
||||
cpu_state_2,
|
||||
generator_transformer_2.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict_2, weight_path_2)
|
||||
|
||||
logger.info(
|
||||
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
|
||||
rank,
|
||||
weight_path_2,
|
||||
local_main_process_only=False)
|
||||
logger.info(
|
||||
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
|
||||
rank,
|
||||
weight_path_2,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Save model config
|
||||
config_dict_2 = generator_transformer_2.hf_config
|
||||
if "dtype" in config_dict_2:
|
||||
del config_dict_2["dtype"] # TODO
|
||||
config_path_2 = os.path.join(inference_save_dir_2,
|
||||
"config.json")
|
||||
with open(config_path_2, "w") as f:
|
||||
json.dump(config_dict_2, f, indent=4)
|
||||
logger.info(
|
||||
"--> generator_2 distillation checkpoint saved at step %s to %s",
|
||||
step, weight_path_2)
|
||||
# Save model config
|
||||
config_dict_2 = generator_transformer_2.hf_config
|
||||
if "dtype" in config_dict_2:
|
||||
del config_dict_2["dtype"] # TODO
|
||||
config_path_2 = os.path.join(inference_save_dir_2, "config.json")
|
||||
with open(config_path_2, "w") as f:
|
||||
json.dump(config_dict_2, f, indent=4)
|
||||
logger.info(
|
||||
"--> generator_2 distillation checkpoint saved at step %s to %s",
|
||||
step, weight_path_2)
|
||||
|
||||
|
||||
def load_checkpoint(transformer,
|
||||
@@ -644,18 +700,37 @@ def load_distillation_checkpoint(
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Load EMA state if available and generator_ema is provided
|
||||
# Load EMA separately if saved in rank0_full mode
|
||||
if generator_ema is not None:
|
||||
try:
|
||||
ema_state = generator_states.get("ema")
|
||||
if ema_state is not None:
|
||||
generator_ema.load_state_dict(ema_state)
|
||||
logger.info("rank: %s, generator EMA state loaded successfully",
|
||||
rank)
|
||||
else:
|
||||
logger.info("rank: %s, no EMA state found in checkpoint", rank)
|
||||
if getattr(generator_ema, "mode", None) == "rank0_full":
|
||||
ema_path = os.path.join(checkpoint_path, "ema",
|
||||
"generator_ema.pt")
|
||||
if rank == 0 and os.path.exists(ema_path):
|
||||
ema_state = torch.load(ema_path, map_location="cpu")
|
||||
generator_ema.load_state_dict(ema_state)
|
||||
logger.info(
|
||||
"rank: %s, generator EMA (rank0_full) loaded from %s",
|
||||
rank, ema_path)
|
||||
elif rank == 0:
|
||||
logger.info(
|
||||
"rank: %s, generator EMA file not found at %s; skipping",
|
||||
rank, ema_path)
|
||||
elif getattr(generator_ema, "mode", None) == "local_shard":
|
||||
ema_path = os.path.join(checkpoint_path, "ema_local_shard",
|
||||
f"generator_ema_rank{rank}.pt")
|
||||
if os.path.exists(ema_path):
|
||||
ema_state = torch.load(ema_path, map_location="cpu")
|
||||
generator_ema.load_state_dict(ema_state)
|
||||
logger.info(
|
||||
"rank: %s, generator EMA shard (local_shard) loaded from %s",
|
||||
rank, ema_path)
|
||||
else:
|
||||
logger.info(
|
||||
"rank: %s, generator EMA shard file not found at %s; skipping",
|
||||
rank, ema_path)
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed to load EMA state: %s", rank,
|
||||
logger.warning("rank: %s, failed to load generator EMA: %s", rank,
|
||||
str(e))
|
||||
|
||||
# Load generator_2 distributed checkpoint (MoE support)
|
||||
@@ -850,7 +925,7 @@ def load_distillation_checkpoint(
|
||||
|
||||
def normalize_dit_input(model_type, latents, vae) -> torch.Tensor:
|
||||
if model_type == "hunyuan_hf" or model_type == "hunyuan":
|
||||
return latents * 0.476986
|
||||
return latents * vae.config.scaling_factor
|
||||
elif model_type == "wan":
|
||||
latents_mean = torch.tensor(vae.latents_mean)
|
||||
latents_std = 1.0 / torch.tensor(vae.latents_std)
|
||||
@@ -1169,6 +1244,71 @@ def custom_to_hf_state_dict(
|
||||
return new_state_dict
|
||||
|
||||
|
||||
def _save_full_ema_safetensors_from_state(
|
||||
state_dict: dict[str, Any],
|
||||
reverse_param_names_mapping: dict[str, tuple[str, int, int]],
|
||||
output_path: str,
|
||||
) -> None:
|
||||
"""
|
||||
Convert a training-format state_dict to HF format and save as safetensors.
|
||||
"""
|
||||
diffusers_state_dict = custom_to_hf_state_dict(state_dict,
|
||||
reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict, output_path)
|
||||
|
||||
|
||||
def _save_rank0_full_ema_safetensors(
|
||||
ema: "EMA_FSDP",
|
||||
module,
|
||||
rank: int,
|
||||
save_dir: str,
|
||||
base_name: str,
|
||||
) -> None:
|
||||
if rank != 0:
|
||||
return
|
||||
ema_dir = os.path.join(save_dir, "ema")
|
||||
os.makedirs(ema_dir, exist_ok=True)
|
||||
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
|
||||
ema_state = ema.state_dict()
|
||||
_save_full_ema_safetensors_from_state(ema_state,
|
||||
module.reverse_param_names_mapping,
|
||||
output_path)
|
||||
logger.info("rank: %s, saved %s as consolidated EMA safetensors to %s",
|
||||
rank,
|
||||
base_name,
|
||||
output_path,
|
||||
local_main_process_only=False)
|
||||
|
||||
|
||||
def _consolidate_local_shard_ema_and_save_safetensors(
|
||||
ema: "EMA_FSDP",
|
||||
module,
|
||||
rank: int,
|
||||
save_dir: str,
|
||||
base_name: str,
|
||||
) -> None:
|
||||
try:
|
||||
# Temporarily apply EMA to the live (sharded) module and gather full CPU state on rank 0
|
||||
with ema.apply_to_model(module):
|
||||
cpu_state_full = gather_state_dict_on_cpu_rank0(module, device=None)
|
||||
if rank == 0:
|
||||
ema_dir = os.path.join(save_dir, "ema")
|
||||
os.makedirs(ema_dir, exist_ok=True)
|
||||
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
|
||||
_save_full_ema_safetensors_from_state(
|
||||
cpu_state_full, module.reverse_param_names_mapping, output_path)
|
||||
logger.info(
|
||||
"rank: %s, saved consolidated %s EMA (from local_shard) as safetensors to %s",
|
||||
rank,
|
||||
base_name,
|
||||
output_path,
|
||||
local_main_process_only=False)
|
||||
except Exception as ce:
|
||||
logger.warning(
|
||||
"rank: %s, failed consolidating %s EMA (local_shard): %s", rank,
|
||||
base_name, str(ce))
|
||||
|
||||
|
||||
def shift_timestep(timestep: torch.Tensor, shift: float,
|
||||
num_train_timestep: float) -> torch.Tensor:
|
||||
if shift == 1:
|
||||
@@ -1795,5 +1935,5 @@ class EMA_FSDP:
|
||||
self.saved.clear()
|
||||
return False
|
||||
|
||||
def apply_to_model(self, module):
|
||||
def apply_to_model(self, module: torch.nn.Module) -> _ApplyEMACtx:
|
||||
return EMA_FSDP._ApplyEMACtx(self, module)
|
||||
|
||||
@@ -10,7 +10,10 @@ from fastvideo.pipelines.basic.wan.wan_dmd_pipeline import WanDMDPipeline
|
||||
from fastvideo.training.distillation_pipeline import DistillationPipeline
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -19,7 +19,10 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
|
||||
from fastvideo.training.distillation_pipeline import DistillationPipeline
|
||||
from fastvideo.utils import is_vsa_available, shallow_asdict
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -18,7 +18,10 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.utils import is_vsa_available, shallow_asdict
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -10,7 +10,10 @@ from fastvideo.training.self_forcing_distillation_pipeline import (
|
||||
SelfForcingDistillationPipeline)
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -10,7 +10,10 @@ from fastvideo.pipelines.basic.wan.wan_pipeline import WanPipeline
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -650,8 +650,10 @@ class WorkerMultiprocProc:
|
||||
logging_info = None
|
||||
if envs.FASTVIDEO_STAGE_LOGGING:
|
||||
logging_info = output_batch.logging_info
|
||||
# result tensor shared by CUDA IPC to avoid serialization overhead
|
||||
result = output_batch.output
|
||||
self.pipe.send({
|
||||
"output_batch": output_batch.output.cpu(),
|
||||
"output_batch": result,
|
||||
"logging_info": logging_info,
|
||||
"extra": output_batch.extra,
|
||||
})
|
||||
|
||||
@@ -26,24 +26,18 @@ class PreprocessingDataValidator:
|
||||
def __init__(self,
|
||||
max_height: int = 1024,
|
||||
max_width: int = 1024,
|
||||
max_h_div_w_ratio: float = 17 / 16,
|
||||
min_h_div_w_ratio: float = 8 / 16,
|
||||
num_frames: int = 16,
|
||||
train_fps: int = 24,
|
||||
speed_factor: float = 1.0,
|
||||
video_length_tolerance_range: float = 5.0,
|
||||
drop_short_ratio: float = 0.0,
|
||||
hw_aspect_threshold: float = 1.5):
|
||||
drop_short_ratio: float = 0.0):
|
||||
self.max_height = max_height
|
||||
self.max_width = max_width
|
||||
self.max_h_div_w_ratio = max_h_div_w_ratio
|
||||
self.min_h_div_w_ratio = min_h_div_w_ratio
|
||||
self.num_frames = num_frames
|
||||
self.train_fps = train_fps
|
||||
self.speed_factor = speed_factor
|
||||
self.video_length_tolerance_range = video_length_tolerance_range
|
||||
self.drop_short_ratio = drop_short_ratio
|
||||
self.hw_aspect_threshold = hw_aspect_threshold
|
||||
self.validators: dict[str, Callable[[dict[str, Any]], bool]] = {}
|
||||
self.filter_counts: dict[str, int] = {}
|
||||
|
||||
@@ -86,26 +80,11 @@ class PreprocessingDataValidator:
|
||||
def _validate_resolution(self, batch: dict[str, Any]) -> bool:
|
||||
"""Validate resolution constraints"""
|
||||
|
||||
aspect = self.max_height / self.max_width
|
||||
if batch["resolution"] is not None:
|
||||
height = batch["resolution"].get("height", None)
|
||||
width = batch["resolution"].get("width", None)
|
||||
|
||||
if height is None or width is None:
|
||||
return False
|
||||
|
||||
return self._filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=self.hw_aspect_threshold * aspect,
|
||||
min_h_div_w_ratio=1 / self.hw_aspect_threshold * aspect,
|
||||
)
|
||||
|
||||
def _filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
|
||||
min_h_div_w_ratio: float) -> bool:
|
||||
"""Filter based on aspect ratio"""
|
||||
return (min_h_div_w_ratio <= h / w <= max_h_div_w_ratio) and (
|
||||
self.min_h_div_w_ratio <= h / w <= self.max_h_div_w_ratio)
|
||||
return not (height is None or width is None)
|
||||
|
||||
def _validate_frame_sampling(self, batch: dict[str, Any]) -> bool:
|
||||
"""Validate frame sampling constraints"""
|
||||
|
||||
Reference in New Issue
Block a user