Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ad16289871 | ||
|
|
2a41da1e6b | ||
|
|
32133171da | ||
|
|
19674c6f29 | ||
|
|
508afb7002 | ||
|
|
288ea88105 | ||
|
|
eb0f1318f3 | ||
|
|
ce9b5910cc | ||
|
|
d0e5a6214a |
+10
-20
@@ -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"
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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) @\
|
||||
|
||||
@@ -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)
|
||||
@@ -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",
|
||||
|
||||
@@ -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
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user