Compare commits

...
18 changed files with 192 additions and 66 deletions
+10 -20
View File
@@ -1,13 +1,12 @@
env:
IMAGE_VERSION: "py3.12-latest"
BUILDKITE_CLEAN_CHECKOUT: true
steps:
- label: "pre-commit"
command: ".buildkite/scripts/pre_commit.sh"
agents:
queue: "default"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- wait
@@ -23,10 +22,9 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=encoder
agents:
queue: "default"
@@ -37,10 +35,9 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=vae
agents:
queue: "default"
@@ -53,20 +50,18 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Transformer Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=transformer
agents:
queue: "default"
- path:
- "fastvideo/v1/**/*.py"
config:
command: "timeout 60m .buildkite/scripts/pr_test.sh"
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=ssim
agents:
queue: "default"
@@ -75,10 +70,9 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Training Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=training
agents:
queue: "default"
@@ -92,10 +86,9 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Training Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=training_vsa
agents:
queue: "default"
@@ -108,10 +101,9 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Inference Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=inference_sta
agents:
queue: "default"
@@ -123,10 +115,9 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=precision_sta
agents:
queue: "default"
@@ -139,10 +130,9 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=precision_vsa
agents:
queue: "default"
+1 -1
View File
@@ -58,7 +58,7 @@ if [ -z "${TEST_TYPE:-}" ]; then
fi
log "Test type: $TEST_TYPE"
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT IMAGE_VERSION=$IMAGE_VERSION"
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$BUILDKITE_PULL_REQUEST IMAGE_VERSION=$IMAGE_VERSION"
case "$TEST_TYPE" in
"encoder")
+14
View File
@@ -15,6 +15,11 @@ With FastVideo's optimizations, you can achieve more than 3x inference improveme
<img src=assets/perf.png width="90%"/>
</div>
## NEWS
- ```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:
@@ -128,6 +133,15 @@ We thank MBZUAI and [Anyscale](https://www.anyscale.com/) for their support thro
If you use FastVideo for your research, please cite our paper:
```bibtex
@misc{zhang2025vsafastervideodiffusion,
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
author={Peiyuan Zhang and Haofeng Huang and Yongqi Chen and Will Lin and Zhengzhong Liu and Ion Stoica and Eric Xing and Hao Zhang},
year={2025},
eprint={2505.13389},
archivePrefix={arXiv},
primaryClass={cs.CV},
url={https://arxiv.org/abs/2505.13389},
}
@misc{zhang2025fastvideogenerationsliding,
title={Fast Video Generation with Sliding Tile Attention},
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
@@ -22,4 +22,5 @@ torchrun --nproc_per_node=$GPU_NUM \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "i2v"
@@ -22,4 +22,5 @@ torchrun --nproc_per_node=$GPU_NUM \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
+1 -1
View File
@@ -52,7 +52,7 @@ class PipelineConfig:
# VAE configuration
vae_config: VAEConfig = field(default_factory=VAEConfig)
vae_precision: str = "fp16"
vae_precision: str = "fp32"
vae_tiling: bool = True
vae_sp: bool = True
+1 -1
View File
@@ -50,7 +50,7 @@ class WanT2V480PConfig(PipelineConfig):
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
vae_precision: str = "fp32"
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp32", ))
@@ -4,7 +4,7 @@
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_precision": "fp32",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
@@ -4,7 +4,7 @@
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_precision": "fp32",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
+6 -4
View File
@@ -413,7 +413,7 @@ class TrainingArgs(FastVideoArgs):
lr_scheduler: str = "constant"
lr_warmup_steps: int = 0
max_grad_norm: float = 0.0
gradient_checkpointing: bool = False
enable_gradient_checkpointing_type: Optional[str] = None
selective_checkpointing: float = 0.0
allow_tf32: bool = False
mixed_precision: str = ""
@@ -612,9 +612,11 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--max-grad-norm",
type=float,
help="Maximum gradient norm")
parser.add_argument("--gradient-checkpointing",
action=StoreBoolean,
help="Whether to use gradient checkpointing")
parser.add_argument("--enable-gradient-checkpointing-type",
type=str,
choices=["full", "ops", "block_skip"],
default=None,
help="Gradient checkpointing type")
parser.add_argument("--selective-checkpointing",
type=float,
help="Selective checkpointing threshold")
+7 -6
View File
@@ -71,15 +71,16 @@ class BaseLayerWithLoRA(nn.Module):
f"cuda:{torch.cuda.current_device()}").full_tensor()
data += (self.slice_lora_b_weights(self.lora_B)
@ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight.data = distribute_tensor(
data, mesh, placements=placements).to(current_device)
self.base_layer.weight = nn.Parameter(
distribute_tensor(data, mesh,
placements=placements).to(current_device))
else:
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(
data = self.base_layer.weight.to(
f"cuda:{torch.cuda.current_device()}")
data += \
(self.slice_lora_b_weights(self.lora_B) @ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight.data = data.to(current_device)
self.base_layer.weight = nn.Parameter(data.to(current_device))
self.merged = True
@torch.no_grad()
@@ -106,8 +107,8 @@ class BaseLayerWithLoRA(nn.Module):
f"cuda:{torch.cuda.current_device()}").full_tensor()
data -= self.slice_lora_b_weights(
self.lora_B) @ self.slice_lora_a_weights(self.lora_A)
self.base_layer.weight.data = distribute_tensor(
data, mesh, placements=placement).to(device)
self.base_layer.weight = nn.Parameter(
distribute_tensor(data, mesh, placements=placement).to(device))
else:
self.base_layer.weight.data -= \
self.slice_lora_b_weights(self.lora_B) @\
+24 -11
View File
@@ -18,6 +18,8 @@ image = (
"PATH": "/root/.cargo/bin:$PATH",
"BUILDKITE_REPO": os.environ.get("BUILDKITE_REPO", ""),
"BUILDKITE_COMMIT": os.environ.get("BUILDKITE_COMMIT", ""),
"BUILDKITE_PULL_REQUEST": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
"IMAGE_VERSION": os.environ.get("IMAGE_VERSION", ""),
})
)
@@ -29,16 +31,27 @@ def run_test(pytest_command: str):
git_repo = os.environ.get("BUILDKITE_REPO")
git_commit = os.environ.get("BUILDKITE_COMMIT")
pr_number = os.environ.get("BUILDKITE_PULL_REQUEST")
print(f"Cloning repository: {git_repo}")
print(f"Checking out commit: {git_commit}")
print(f"Target commit: {git_commit}")
if pr_number:
print(f"PR number: {pr_number}")
# For PRs (including forks), use GitHub's PR refs to get the correct commit
if pr_number and pr_number != "false":
checkout_command = f"git fetch --prune origin refs/pull/{pr_number}/head && git checkout FETCH_HEAD"
print(f"Using PR ref for checkout: {checkout_command}")
else:
checkout_command = f"git checkout {git_commit}"
print(f"Using direct commit checkout: {checkout_command}")
command = f"""
source $HOME/.local/bin/env &&
source /opt/venv/bin/activate &&
git clone {git_repo} /FastVideo &&
cd /FastVideo &&
git checkout {git_commit} &&
{checkout_command} &&
uv pip install -e .[test] &&
{pytest_command}
"""
@@ -49,38 +62,38 @@ def run_test(pytest_command: str):
sys.exit(result.returncode)
@app.function(gpu="L40S:1", image=image, timeout=1800)
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_encoder_tests():
run_test("pytest ./fastvideo/v1/tests/encoders -vs")
@app.function(gpu="L40S:1", image=image, timeout=1800)
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_vae_tests():
run_test("pytest ./fastvideo/v1/tests/vaes -vs")
@app.function(gpu="L40S:1", image=image, timeout=1800)
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_transformer_tests():
run_test("pytest ./fastvideo/v1/tests/transformers -vs")
@app.function(gpu="L40S:2", image=image, timeout=3600)
@app.function(gpu="L40S:2", image=image, timeout=1800)
def run_ssim_tests():
run_test("pytest ./fastvideo/v1/tests/ssim -vs")
@app.function(gpu="L40S:4", image=image, timeout=1800, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/v1/tests/training/Vanilla -srP")
@app.function(gpu="H100:2", image=image, timeout=1800, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
@app.function(gpu="H100:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests_VSA():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/v1/tests/training/VSA -srP")
@app.function(gpu="H100:2", image=image, timeout=1800)
@app.function(gpu="H100:2", image=image, timeout=900)
def run_inference_tests_STA():
run_test("pytest ./fastvideo/v1/tests/inference/STA -srP")
@app.function(gpu="H100:1", image=image, timeout=1800)
@app.function(gpu="H100:1", image=image, timeout=900)
def run_precision_tests_STA():
run_test("python csrc/attn/tests/test_sta.py")
@app.function(gpu="H100:1", image=image, timeout=1800)
@app.function(gpu="H100:1", image=image, timeout=900)
def run_precision_tests_VSA():
run_test("python csrc/attn/tests/test_block_sparse.py")
@@ -0,0 +1,91 @@
import collections
from enum import Enum
from typing import Optional
import torch
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
checkpoint_wrapper)
TRANSFORMER_BLOCK_NAMES = [
"blocks",
"double_blocks",
"single_blocks",
"transformer_blocks",
"temporal_transformer_blocks",
"transformer_double_blocks",
"transformer_single_blocks",
]
class CheckpointType(str, Enum):
FULL = "full"
OPS = "ops"
BLOCK_SKIP = "block_skip"
_SELECTIVE_ACTIVATION_CHECKPOINTING_OPS = {
torch.ops.aten.mm.default,
torch.ops.aten._scaled_dot_product_efficient_attention.default,
torch.ops.aten._scaled_dot_product_flash_attention.default,
torch.ops._c10d_functional.reduce_scatter_tensor.default,
}
def apply_activation_checkpointing(
module: torch.nn.Module,
checkpointing_type: str = CheckpointType.FULL,
n_layer: int = 1) -> torch.nn.Module:
if checkpointing_type == CheckpointType.FULL:
module = _apply_activation_checkpointing_blocks(module)
elif checkpointing_type == CheckpointType.OPS:
module = _apply_activation_checkpointing_ops(
module, _SELECTIVE_ACTIVATION_CHECKPOINTING_OPS)
elif checkpointing_type == CheckpointType.BLOCK_SKIP:
module = _apply_activation_checkpointing_blocks(module, n_layer)
else:
raise ValueError(
f"Checkpointing type '{checkpointing_type}' not supported. Supported types are {CheckpointType.__members__.keys()}"
)
return module
def _apply_activation_checkpointing_blocks(
module: torch.nn.Module,
n_layer: Optional[int] = None) -> torch.nn.Module:
for transformer_block_name in TRANSFORMER_BLOCK_NAMES:
blocks: torch.nn.Module = getattr(module, transformer_block_name, None)
if blocks is None:
continue
for index, (layer_id, block) in enumerate(blocks.named_children()):
if n_layer is None or index % n_layer == 0:
block = checkpoint_wrapper(block, preserve_rng_state=False)
blocks.register_module(layer_id, block)
return module
def _apply_activation_checkpointing_ops(module: torch.nn.Module,
ops) -> torch.nn.Module:
from torch.utils.checkpoint import (CheckpointPolicy,
create_selective_checkpoint_contexts)
def _get_custom_policy(meta: dict[str, int]) -> CheckpointPolicy:
def _custom_policy(ctx, func, *args, **kwargs):
mode = "recompute" if ctx.is_recompute else "forward"
mm_count_key = f"{mode}_mm_count"
if func == torch.ops.aten.mm.default:
meta[mm_count_key] += 1
# Saves output of all compute ops, except every second mm
to_save = func in ops and not (func == torch.ops.aten.mm.default
and meta[mm_count_key] % 2 == 0)
return CheckpointPolicy.MUST_SAVE if to_save else CheckpointPolicy.PREFER_RECOMPUTE
return _custom_policy
def selective_checkpointing_context_fn():
meta: dict[str, int] = collections.defaultdict(int)
return create_selective_checkpoint_contexts(_get_custom_policy(meta))
return checkpoint_wrapper(module,
context_fn=selective_checkpointing_context_fn,
preserve_rng_state=False)
+23 -11
View File
@@ -31,6 +31,8 @@ from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
TrainingBatch)
from fastvideo.v1.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.v1.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
@@ -82,6 +84,11 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self.transformer.requires_grad_(True)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
self.transformer = apply_activation_checkpointing(
self.transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
noise_scheduler = self.modules["scheduler"]
params_to_optimize = self.transformer.parameters()
@@ -309,17 +316,18 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
model_pred = self.transformer(**input_kwargs)
if self.training_args.precondition_outputs:
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
if self.training_args.precondition_outputs:
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
# make sure no implicit broadcasting happens
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
loss = (torch.mean((model_pred.float() - target.float())**2) /
self.training_args.gradient_accumulation_steps)
# make sure no implicit broadcasting happens
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
loss = (torch.mean((model_pred.float() - target.float())**2) /
self.training_args.gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
loss.backward()
avg_loss = loss.detach().clone()
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
# local_main_process_only=False)
world_group = get_world_group()
@@ -546,6 +554,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
negative_prompt_attention_mask: torch.Tensor | None
) -> ForwardBatch:
assert len(validation_batch['info_list']
) == 1, "Only batch size 1 is supported for validation"
prompt = validation_batch['info_list'][0]['prompt']
prompt_embeds = validation_batch['text_embedding']
prompt_attention_mask = validation_batch['text_attention_mask']
@@ -629,7 +639,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Process each validation prompt for each validation step
for num_inference_steps in validation_steps:
step_videos: List[np.ndarray] = []
step_captions: List[str | None] = []
step_captions: List[str] = []
for validation_batch in validation_dataloader:
batch = self._prepare_validation_inputs(
@@ -637,7 +647,9 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
num_inference_steps, negative_prompt_embeds,
negative_prompt_attention_mask)
step_captions.extend([None]) # TODO(peiyuan): add caption
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
# Run validation inference
with torch.no_grad(), torch.autocast("cuda",
+5 -5
View File
@@ -162,8 +162,8 @@ def save_checkpoint(transformer,
weight_path,
local_main_process_only=False)
# Convert training format to diffusers format and save
diffusers_state_dict = convert_training_to_diffusers_format(
# Convert fastvideo custom format to diffusers format and save
diffusers_state_dict = convert_custom_format_to_diffusers_format(
cpu_state, transformer)
save_file(diffusers_state_dict, weight_path)
@@ -487,10 +487,10 @@ def _has_foreach_support(tensors: List[torch.Tensor],
t is None or type(t) in [torch.Tensor] for t in tensors)
def convert_training_to_diffusers_format(state_dict: Dict[str, Any],
transformer) -> Dict[str, Any]:
def convert_custom_format_to_diffusers_format(state_dict: Dict[str, Any],
transformer) -> Dict[str, Any]:
"""
Convert training format state dict to diffusers format using reverse_param_names_mapping.
Convert fastvideo custom format state dict to diffusers format using reverse_param_names_mapping.
Args:
state_dict: State dict in training format
+1 -2
View File
@@ -1,11 +1,10 @@
# trigger test
[build-system]
requires = ["setuptools>=61.0"]
build-backend = "setuptools.build_meta"
[project]
name = "fastvideo"
version = "0.1.0"
version = "0.1.1"
description = "FastVideo"
readme = "README.md"
requires-python = ">=3.8"
+2 -1
View File
@@ -48,4 +48,5 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--weight_decay 0.01 \
--not_apply_cfg_solver \
--dit_precision "fp32" \
--max_grad_norm 1.0
--max_grad_norm 1.0 \
--enable_gradient_checkpointing_type "full"
+2 -1
View File
@@ -58,5 +58,6 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--VSA_decay_sparsity 0.9 \
--VSA_decay_rate 0.03 \
--VSA_decay_interval_steps 30 \
--VSA_val_sparsity 0.9
--VSA_val_sparsity 0.9 \
--enable_gradient_checkpointing_type "full"
# --resume_from_checkpoint "$CHECKPOINT_PATH"