Compare commits

...
27 changed files with 482 additions and 310 deletions
+18 -20
View File
@@ -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])
+24
View File
@@ -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:
+1 -70
View File
@@ -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)
+41 -2
View File
@@ -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
+5 -17
View File
@@ -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
View File
@@ -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)
+3 -2
View File
@@ -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):
+1 -1
View File
@@ -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(
+62 -24
View File
@@ -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:
+2 -2
View File
@@ -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(
+4 -3
View File
@@ -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]
+1 -1
View File
@@ -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
+5 -2
View File
@@ -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__)
+13 -26
View File
@@ -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):
+47 -34
View File
@@ -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
+191 -51
View File
@@ -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__)
+4 -1
View File
@@ -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__)
+3 -1
View File
@@ -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,
})
+2 -23
View File
@@ -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"""