Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
80d5b16e1b | ||
|
|
7dcfe4ea8f | ||
|
|
9ab7f031af |
+20
-10
@@ -1,12 +1,13 @@
|
||||
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
|
||||
|
||||
@@ -22,9 +23,10 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Encoder Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=encoder
|
||||
agents:
|
||||
queue: "default"
|
||||
@@ -35,9 +37,10 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "VAE Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=vae
|
||||
agents:
|
||||
queue: "default"
|
||||
@@ -50,18 +53,20 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 30m .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 30m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 60m .buildkite/scripts/pr_test.sh"
|
||||
label: "SSIM Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=ssim
|
||||
agents:
|
||||
queue: "default"
|
||||
@@ -70,9 +75,10 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Training Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=training
|
||||
agents:
|
||||
queue: "default"
|
||||
@@ -86,9 +92,10 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Training Tests VSA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=training_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
@@ -101,9 +108,10 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Inference Tests STA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=inference_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
@@ -115,9 +123,10 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests STA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=precision_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
@@ -130,9 +139,10 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests VSA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=precision_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
@@ -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 BUILDKITE_PULL_REQUEST=$BUILDKITE_PULL_REQUEST IMAGE_VERSION=$IMAGE_VERSION"
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
|
||||
@@ -15,11 +15,6 @@ 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:
|
||||
@@ -133,15 +128,6 @@ 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,5 +22,4 @@ 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,5 +22,4 @@ 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"
|
||||
@@ -52,7 +52,7 @@ class PipelineConfig:
|
||||
|
||||
# VAE configuration
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
vae_precision: str = "fp32"
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = True
|
||||
|
||||
|
||||
@@ -50,7 +50,7 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp32"
|
||||
vae_precision: str = "fp16"
|
||||
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": "fp32",
|
||||
"vae_precision": "fp16",
|
||||
"vae_tiling": false,
|
||||
"vae_sp": false,
|
||||
"vae_config": {
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
"use_cpu_offload": true,
|
||||
"disable_autocast": false,
|
||||
"precision": "bf16",
|
||||
"vae_precision": "fp32",
|
||||
"vae_precision": "fp16",
|
||||
"vae_tiling": false,
|
||||
"vae_sp": false,
|
||||
"vae_config": {
|
||||
|
||||
@@ -413,7 +413,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
lr_scheduler: str = "constant"
|
||||
lr_warmup_steps: int = 0
|
||||
max_grad_norm: float = 0.0
|
||||
enable_gradient_checkpointing_type: Optional[str] = None
|
||||
gradient_checkpointing: bool = False
|
||||
selective_checkpointing: float = 0.0
|
||||
allow_tf32: bool = False
|
||||
mixed_precision: str = ""
|
||||
@@ -612,11 +612,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--max-grad-norm",
|
||||
type=float,
|
||||
help="Maximum gradient norm")
|
||||
parser.add_argument("--enable-gradient-checkpointing-type",
|
||||
type=str,
|
||||
choices=["full", "ops", "block_skip"],
|
||||
default=None,
|
||||
help="Gradient checkpointing type")
|
||||
parser.add_argument("--gradient-checkpointing",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use gradient checkpointing")
|
||||
parser.add_argument("--selective-checkpointing",
|
||||
type=float,
|
||||
help="Selective checkpointing threshold")
|
||||
|
||||
@@ -71,16 +71,15 @@ 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 = nn.Parameter(
|
||||
distribute_tensor(data, mesh,
|
||||
placements=placements).to(current_device))
|
||||
self.base_layer.weight.data = distribute_tensor(
|
||||
data, mesh, placements=placements).to(current_device)
|
||||
else:
|
||||
current_device = self.base_layer.weight.data.device
|
||||
data = self.base_layer.weight.to(
|
||||
data = self.base_layer.weight.data.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 = nn.Parameter(data.to(current_device))
|
||||
self.base_layer.weight.data = data.to(current_device)
|
||||
self.merged = True
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -107,8 +106,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 = nn.Parameter(
|
||||
distribute_tensor(data, mesh, placements=placement).to(device))
|
||||
self.base_layer.weight.data = distribute_tensor(
|
||||
data, mesh, placements=placement).to(device)
|
||||
else:
|
||||
self.base_layer.weight.data -= \
|
||||
self.slice_lora_b_weights(self.lora_B) @\
|
||||
|
||||
@@ -18,8 +18,6 @@ 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", ""),
|
||||
})
|
||||
)
|
||||
|
||||
@@ -31,27 +29,16 @@ 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"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}")
|
||||
print(f"Checking out commit: {git_commit}")
|
||||
|
||||
command = f"""
|
||||
source $HOME/.local/bin/env &&
|
||||
source /opt/venv/bin/activate &&
|
||||
git clone {git_repo} /FastVideo &&
|
||||
cd /FastVideo &&
|
||||
{checkout_command} &&
|
||||
git checkout {git_commit} &&
|
||||
uv pip install -e .[test] &&
|
||||
{pytest_command}
|
||||
"""
|
||||
@@ -62,38 +49,38 @@ def run_test(pytest_command: str):
|
||||
|
||||
sys.exit(result.returncode)
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1800)
|
||||
def run_encoder_tests():
|
||||
run_test("pytest ./fastvideo/v1/tests/encoders -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1800)
|
||||
def run_vae_tests():
|
||||
run_test("pytest ./fastvideo/v1/tests/vaes -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1800)
|
||||
def run_transformer_tests():
|
||||
run_test("pytest ./fastvideo/v1/tests/transformers -vs")
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=1800)
|
||||
@app.function(gpu="L40S:2", image=image, timeout=3600)
|
||||
def run_ssim_tests():
|
||||
run_test("pytest ./fastvideo/v1/tests/ssim -vs")
|
||||
|
||||
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
@app.function(gpu="L40S:4", image=image, timeout=1800, 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=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
@app.function(gpu="H100:2", image=image, timeout=1800, 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=900)
|
||||
@app.function(gpu="H100:2", image=image, timeout=1800)
|
||||
def run_inference_tests_STA():
|
||||
run_test("pytest ./fastvideo/v1/tests/inference/STA -srP")
|
||||
|
||||
@app.function(gpu="H100:1", image=image, timeout=900)
|
||||
@app.function(gpu="H100:1", image=image, timeout=1800)
|
||||
def run_precision_tests_STA():
|
||||
run_test("python csrc/attn/tests/test_sta.py")
|
||||
|
||||
@app.function(gpu="H100:1", image=image, timeout=900)
|
||||
@app.function(gpu="H100:1", image=image, timeout=1800)
|
||||
def run_precision_tests_VSA():
|
||||
run_test("python csrc/attn/tests/test_block_sparse.py")
|
||||
|
||||
@@ -1,91 +0,0 @@
|
||||
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)
|
||||
@@ -31,8 +31,6 @@ 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,
|
||||
@@ -84,11 +82,6 @@ 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()
|
||||
@@ -316,18 +309,17 @@ 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)
|
||||
|
||||
loss.backward()
|
||||
avg_loss = loss.detach().clone()
|
||||
# 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()
|
||||
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
|
||||
# local_main_process_only=False)
|
||||
world_group = get_world_group()
|
||||
@@ -554,8 +546,6 @@ 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']
|
||||
@@ -639,7 +629,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] = []
|
||||
step_captions: List[str | None] = []
|
||||
|
||||
for validation_batch in validation_dataloader:
|
||||
batch = self._prepare_validation_inputs(
|
||||
@@ -647,9 +637,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
num_inference_steps, negative_prompt_embeds,
|
||||
negative_prompt_attention_mask)
|
||||
|
||||
assert batch.prompt is not None and isinstance(
|
||||
batch.prompt, str)
|
||||
step_captions.append(batch.prompt)
|
||||
step_captions.extend([None]) # TODO(peiyuan): add caption
|
||||
|
||||
# Run validation inference
|
||||
with torch.no_grad(), torch.autocast("cuda",
|
||||
|
||||
@@ -162,8 +162,8 @@ def save_checkpoint(transformer,
|
||||
weight_path,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Convert fastvideo custom format to diffusers format and save
|
||||
diffusers_state_dict = convert_custom_format_to_diffusers_format(
|
||||
# Convert training format to diffusers format and save
|
||||
diffusers_state_dict = convert_training_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_custom_format_to_diffusers_format(state_dict: Dict[str, Any],
|
||||
transformer) -> Dict[str, Any]:
|
||||
def convert_training_to_diffusers_format(state_dict: Dict[str, Any],
|
||||
transformer) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert fastvideo custom format state dict to diffusers format using reverse_param_names_mapping.
|
||||
Convert training format state dict to diffusers format using reverse_param_names_mapping.
|
||||
|
||||
Args:
|
||||
state_dict: State dict in training format
|
||||
|
||||
@@ -4,11 +4,12 @@ from copy import deepcopy
|
||||
from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.dataset.dataloader.schema import (
|
||||
pyarrow_schema_i2v, pyarrow_schema_i2v_validation)
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_torch_device, get_sp_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
@@ -18,6 +19,8 @@ from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
|
||||
from fastvideo.v1.pipelines.wan.wan_i2v_pipeline import (
|
||||
WanImageToVideoValidationPipeline)
|
||||
from fastvideo.v1.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.v1.training.training_utils import (shard_latents_across_sp,
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases)
|
||||
from fastvideo.v1.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
@@ -204,6 +207,36 @@ class WanI2VTrainingPipeline(TrainingPipeline):
|
||||
)
|
||||
return batch
|
||||
|
||||
def _clip_grad_norm(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Override to add gradient synchronization across SP ranks."""
|
||||
assert self.training_args is not None
|
||||
max_grad_norm = self.training_args.max_grad_norm
|
||||
|
||||
# CRITICAL FIX: Synchronize gradients across SP ranks before clipping
|
||||
# Different SP ranks compute different gradients due to different noise patterns
|
||||
# These gradients must be averaged across SP ranks for stable training
|
||||
if self.training_args.sp_size > 1:
|
||||
sp_group = get_sp_group()
|
||||
for param in self.transformer.parameters():
|
||||
if param.grad is not None:
|
||||
# Average gradients across SP ranks
|
||||
sp_group.all_reduce(param.grad, op=torch.distributed.ReduceOp.AVG)
|
||||
|
||||
if max_grad_norm is not None:
|
||||
model_parts = [self.transformer]
|
||||
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
|
||||
[p for m in model_parts for p in m.parameters()],
|
||||
max_grad_norm,
|
||||
foreach=None,
|
||||
)
|
||||
assert grad_norm is not float('nan') or grad_norm is not float(
|
||||
'inf')
|
||||
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
|
||||
else:
|
||||
grad_norm = 0.0
|
||||
training_batch.grad_norm = grad_norm
|
||||
return training_batch
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting training pipeline...")
|
||||
|
||||
+2
-1
@@ -1,10 +1,11 @@
|
||||
# trigger test
|
||||
[build-system]
|
||||
requires = ["setuptools>=61.0"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo"
|
||||
version = "0.1.1"
|
||||
version = "0.1.0"
|
||||
description = "FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.8"
|
||||
|
||||
@@ -48,5 +48,4 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
|
||||
--weight_decay 0.01 \
|
||||
--not_apply_cfg_solver \
|
||||
--dit_precision "fp32" \
|
||||
--max_grad_norm 1.0 \
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--max_grad_norm 1.0
|
||||
@@ -58,6 +58,5 @@ 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 \
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--VSA_val_sparsity 0.9
|
||||
# --resume_from_checkpoint "$CHECKPOINT_PATH"
|
||||
|
||||
Reference in New Issue
Block a user