Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ea8c154176 |
+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")
|
||||
|
||||
@@ -122,17 +122,17 @@ jobs:
|
||||
# Actual tests
|
||||
encoder-test:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
- *common-paths
|
||||
vae-test:
|
||||
- 'fastvideo/v1/models/vaes/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/vaes/**'
|
||||
- *common-paths
|
||||
transformer-test:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
@@ -264,7 +264,7 @@ jobs:
|
||||
with:
|
||||
job_id: "training-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 2
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
@@ -284,7 +284,7 @@ jobs:
|
||||
with:
|
||||
job_id: "inference-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 2
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
|
||||
@@ -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},
|
||||
|
||||
@@ -10,7 +10,7 @@ def main():
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# if num_gpus > 1, FastVideo will automatically handle distributed setup
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
use_cpu_offload=False
|
||||
|
||||
@@ -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"
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import Any, Dict
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
@@ -12,7 +12,9 @@ logger = init_logger(__name__)
|
||||
# 3. Any field in ArchConfig is fixed upon initialization, and should be hidden away from users
|
||||
@dataclass
|
||||
class ArchConfig:
|
||||
pass
|
||||
stacked_params_mapping: List[Tuple[str, str, str]] = field(
|
||||
default_factory=list
|
||||
) # mapping from huggingface weight names to custom names
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -5,13 +5,11 @@ from typing import List, Optional, Tuple, Union
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_blocks(n: str, m) -> bool:
|
||||
return "blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class StepVideoArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda:
|
||||
[lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit()])
|
||||
|
||||
_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
|
||||
@@ -32,8 +32,11 @@ class TextEncoderArchConfig(EncoderArchConfig):
|
||||
output_past: bool = True
|
||||
scalable_attention: bool = True
|
||||
tie_word_embeddings: bool = False
|
||||
|
||||
stacked_params_mapping: List[Tuple[str, str, str]] = field(
|
||||
default_factory=list
|
||||
) # mapping from huggingface weight names to custom names
|
||||
tokenizer_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.tokenizer_kwargs = {
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
|
||||
ImageEncoderConfig,
|
||||
@@ -8,6 +8,14 @@ from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
return "layers" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def _is_embeddings(n: str, m) -> bool:
|
||||
return n.endswith("embeddings")
|
||||
|
||||
|
||||
@dataclass
|
||||
class CLIPTextArchConfig(TextEncoderArchConfig):
|
||||
vocab_size: int = 49408
|
||||
@@ -27,6 +35,15 @@ class CLIPTextArchConfig(TextEncoderArchConfig):
|
||||
bos_token_id: int = 49406
|
||||
eos_token_id: int = 49407
|
||||
text_len: int = 77
|
||||
stacked_params_mapping: List[Tuple[str, str,
|
||||
str]] = field(default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [_is_transformer_layer, _is_embeddings])
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -1,11 +1,23 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
return "layers" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def _is_embeddings(n: str, m) -> bool:
|
||||
return n.endswith("embed_tokens")
|
||||
|
||||
|
||||
def _is_final_norm(n: str, m) -> bool:
|
||||
return n.endswith("norm")
|
||||
|
||||
|
||||
@dataclass
|
||||
class LlamaArchConfig(TextEncoderArchConfig):
|
||||
vocab_size: int = 32000
|
||||
@@ -32,6 +44,18 @@ class LlamaArchConfig(TextEncoderArchConfig):
|
||||
head_dim: Optional[int] = None
|
||||
hidden_state_skip_layer: int = 2
|
||||
text_len: int = 256
|
||||
stacked_params_mapping: List[Tuple[str, str, str]] = field(
|
||||
default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q_proj", "q"),
|
||||
(".qkv_proj", ".k_proj", "k"),
|
||||
(".qkv_proj", ".v_proj", "v"),
|
||||
(".gate_up_proj", ".gate_proj", 0), # type: ignore
|
||||
(".gate_up_proj", ".up_proj", 1), # type: ignore
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda:
|
||||
[_is_transformer_layer, _is_embeddings, _is_final_norm])
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -1,11 +1,23 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
return "block" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def _is_embeddings(n: str, m) -> bool:
|
||||
return n.endswith("shared")
|
||||
|
||||
|
||||
def _is_final_layernorm(n: str, m) -> bool:
|
||||
return n.endswith("final_layer_norm")
|
||||
|
||||
|
||||
@dataclass
|
||||
class T5ArchConfig(TextEncoderArchConfig):
|
||||
vocab_size: int = 32128
|
||||
@@ -29,6 +41,16 @@ class T5ArchConfig(TextEncoderArchConfig):
|
||||
eos_token_id: int = 1
|
||||
classifier_dropout: float = 0.0
|
||||
text_len: int = 512
|
||||
stacked_params_mapping: List[Tuple[str, str,
|
||||
str]] = field(default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q", "q"),
|
||||
(".qkv_proj", ".k", "k"),
|
||||
(".qkv_proj", ".v", "v"),
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda:
|
||||
[_is_transformer_layer, _is_embeddings, _is_final_layernorm])
|
||||
|
||||
# Referenced from https://github.com/huggingface/transformers/blob/main/src/transformers/models/t5/configuration_t5.py
|
||||
def __post_init__(self):
|
||||
|
||||
@@ -11,7 +11,7 @@ from fastvideo.v1.dataset.parquet_dataset_iterable_style import (
|
||||
build_parquet_iterable_style_dataloader)
|
||||
from fastvideo.v1.distributed import get_world_rank
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_torch_device,
|
||||
cleanup_dist_env_and_memory, get_local_torch_device,
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
@@ -148,8 +148,8 @@ def main() -> None:
|
||||
break
|
||||
|
||||
# Move data to device
|
||||
latents = latents.to(get_torch_device())
|
||||
embeddings = embeddings.to(get_torch_device())
|
||||
latents = latents.to(get_local_torch_device())
|
||||
embeddings = embeddings.to(get_local_torch_device())
|
||||
|
||||
# Calculate actual batch size
|
||||
batch_size = latents.size(0)
|
||||
|
||||
@@ -13,7 +13,7 @@ from fastvideo.v1.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.v1.distributed import get_world_rank
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_torch_device,
|
||||
cleanup_dist_env_and_memory, get_local_torch_device,
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
@@ -165,8 +165,8 @@ def main() -> None:
|
||||
break
|
||||
|
||||
# Move data to device
|
||||
latents = latents.to(get_torch_device())
|
||||
embeddings = embeddings.to(get_torch_device())
|
||||
latents = latents.to(get_local_torch_device())
|
||||
embeddings = embeddings.to(get_local_torch_device())
|
||||
|
||||
# Calculate actual batch size
|
||||
batch_size = latents.size(0)
|
||||
|
||||
@@ -3,10 +3,10 @@
|
||||
from fastvideo.v1.distributed.communication_op import *
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_dp_group, get_dp_rank, get_dp_world_size,
|
||||
get_sp_group, get_sp_parallel_rank, get_sp_world_size, get_torch_device,
|
||||
get_tp_group, get_tp_rank, get_tp_world_size, get_world_group,
|
||||
get_world_rank, get_world_size, init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
get_local_torch_device, get_sp_group, get_sp_parallel_rank,
|
||||
get_sp_world_size, get_tp_group, get_tp_rank, get_tp_world_size,
|
||||
get_world_group, get_world_rank, get_world_size,
|
||||
init_distributed_environment, initialize_model_parallel,
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
model_parallel_is_initialized)
|
||||
from fastvideo.v1.distributed.utils import *
|
||||
@@ -40,5 +40,5 @@ __all__ = [
|
||||
"get_tp_world_size",
|
||||
|
||||
# Get torch device
|
||||
"get_torch_device",
|
||||
"get_local_torch_device",
|
||||
]
|
||||
|
||||
@@ -36,6 +36,7 @@ from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import Backend, ProcessGroup, ReduceOp
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
@@ -692,6 +693,7 @@ class GroupCoordinator:
|
||||
|
||||
|
||||
_WORLD: Optional[GroupCoordinator] = None
|
||||
_NODE: Optional[GroupCoordinator] = None
|
||||
|
||||
|
||||
def get_world_group() -> GroupCoordinator:
|
||||
@@ -699,6 +701,11 @@ def get_world_group() -> GroupCoordinator:
|
||||
return _WORLD
|
||||
|
||||
|
||||
def get_node_group() -> GroupCoordinator:
|
||||
assert _NODE is not None, ("node group is not initialized")
|
||||
return _NODE
|
||||
|
||||
|
||||
def init_world_group(ranks: List[int], local_rank: int,
|
||||
backend: str) -> GroupCoordinator:
|
||||
return GroupCoordinator(
|
||||
@@ -710,6 +717,18 @@ def init_world_group(ranks: List[int], local_rank: int,
|
||||
)
|
||||
|
||||
|
||||
def init_node_group(local_rank: int, backend: str):
|
||||
cpu_group = get_world_group().cpu_group
|
||||
node_ranks = same_node_ranks(cpu_group)
|
||||
node_size = len(node_ranks)
|
||||
all_node_ranks = [
|
||||
list(range(i * node_size, (i + 1) * node_size))
|
||||
for i in range(dist.get_world_size() // node_size)
|
||||
]
|
||||
global _NODE
|
||||
_NODE = init_model_parallel_group(all_node_ranks, local_rank, backend)
|
||||
|
||||
|
||||
def init_model_parallel_group(
|
||||
group_ranks: List[List[int]],
|
||||
local_rank: int,
|
||||
@@ -782,6 +801,8 @@ def init_distributed_environment(
|
||||
else:
|
||||
assert _WORLD.world_size == torch.distributed.get_world_size(), (
|
||||
"world group already initialized with a different world size")
|
||||
# Init a group for each node
|
||||
init_node_group(local_rank, backend)
|
||||
|
||||
|
||||
_SP: Optional[GroupCoordinator] = None
|
||||
@@ -904,7 +925,7 @@ def get_dp_rank() -> int:
|
||||
return get_dp_group().rank_in_group
|
||||
|
||||
|
||||
def get_torch_device() -> torch.device:
|
||||
def get_local_torch_device() -> torch.device:
|
||||
"""Return the torch device for the current rank."""
|
||||
return torch.device(f"cuda:{envs.LOCAL_RANK}")
|
||||
|
||||
@@ -1021,17 +1042,22 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
|
||||
"torch._C._host_emptyCache() only available in Pytorch >=2.5")
|
||||
|
||||
|
||||
def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
|
||||
source_rank: int = 0) -> List[bool]:
|
||||
def same_node_ranks(pg: Union[ProcessGroup, StatelessProcessGroup],
|
||||
source_rank: int = 0) -> List[int]:
|
||||
"""
|
||||
This is a collective operation that returns if each rank is in the same node
|
||||
This is a collective operation that returns ranks that are in the same node
|
||||
as the source rank. It tests if processes are attached to the same
|
||||
memory system (shared access to shared memory).
|
||||
Args:
|
||||
pg: the global process group to test
|
||||
source_rank: the rank to test against
|
||||
Returns:
|
||||
A list of ranks that are in the same node as the source rank.
|
||||
"""
|
||||
if isinstance(pg, ProcessGroup):
|
||||
assert torch.distributed.get_backend(
|
||||
pg) != torch.distributed.Backend.NCCL, (
|
||||
"in_the_same_node_as should be tested with a non-NCCL group.")
|
||||
"same_node_ranks should be tested with a non-NCCL group.")
|
||||
# local rank inside the group
|
||||
rank = torch.distributed.get_rank(group=pg)
|
||||
world_size = torch.distributed.get_world_size(group=pg)
|
||||
@@ -1103,7 +1129,7 @@ def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
|
||||
rank_data = pg.broadcast_obj(is_in_the_same_node, src=i)
|
||||
aggregated_data += rank_data
|
||||
|
||||
return [x == 1 for x in aggregated_data.tolist()]
|
||||
return [i for i, x in enumerate(aggregated_data.tolist()) if x == 1]
|
||||
|
||||
|
||||
def initialize_tensor_parallel_group(
|
||||
|
||||
@@ -58,8 +58,10 @@ class FastVideoArgs:
|
||||
|
||||
output_type: str = "pil"
|
||||
|
||||
use_cpu_offload: bool = True
|
||||
use_cpu_offload: bool = True # For DiT
|
||||
use_fsdp_inference: bool = True
|
||||
text_encoder_offload: bool = True
|
||||
pin_cpu_memory: bool = True
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
@@ -208,7 +210,7 @@ class FastVideoArgs:
|
||||
"--use-cpu-offload",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Use CPU offload for model inference. Enable if run out of memory with FSDP.",
|
||||
"Use CPU offload for DiT inference. Enable if run out of memory with FSDP.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-fsdp-inference",
|
||||
@@ -216,7 +218,19 @@ class FastVideoArgs:
|
||||
help=
|
||||
"Use FSDP for inference by sharding the model weights. Latency is very low due to prefetch--enable if run out of memory.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--text-encoder-cpu-offload",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Use CPU offload for text encoder. Enable if run out of memory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pin-cpu-memory",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Pin memory for CPU offload. Only added as a temp workaround if it throws \"CUDA error: invalid argument\". "
|
||||
"Should be enabled in almost all cases",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-autocast",
|
||||
action=StoreBoolean,
|
||||
@@ -413,7 +427,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 +626,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")
|
||||
|
||||
@@ -6,6 +6,7 @@ from typing import Optional, Tuple, Union
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.distributed.tensor import DTensor
|
||||
|
||||
from fastvideo.v1.layers.custom_op import CustomOp
|
||||
|
||||
@@ -38,12 +39,6 @@ class RMSNorm(CustomOp):
|
||||
if self.has_weight:
|
||||
self.weight = nn.Parameter(self.weight)
|
||||
|
||||
# if we do fully_shard(model.layer_norm), and we call layer_form.forward_native(input) instead of layer_norm(input),
|
||||
# we need to call model.layer_norm.register_fsdp_forward_method(model, "forward_native") to make sure fsdp2 hooks are triggered
|
||||
# for mixed precision and cpu offloading
|
||||
|
||||
# the even better way might be fully_shard(model.layer_norm, mp_policy=, cpu_offloading=), and call model.layer_norm(input). everything should work out of the box
|
||||
# because fsdp2 hooks will be triggered with model.layer_norm.__call__
|
||||
def forward_native(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
@@ -76,7 +71,12 @@ class RMSNorm(CustomOp):
|
||||
x = x * torch.rsqrt(variance + self.variance_epsilon)
|
||||
x = x.to(orig_dtype)
|
||||
if self.has_weight:
|
||||
x = x * self.weight
|
||||
# TODO(wenxuan): When using CPU offload, FSDP has a bug that doesn't unwrap DTensor in final_layer_norm.
|
||||
# Report this
|
||||
if isinstance(self.weight, DTensor):
|
||||
x = x * self.weight.to(x.device).full_tensor()
|
||||
else:
|
||||
x = x * self.weight
|
||||
if residual is None:
|
||||
return x
|
||||
else:
|
||||
|
||||
@@ -455,10 +455,7 @@ class StepVideoTransformerBlock(nn.Module):
|
||||
|
||||
class StepVideoModel(BaseDiT):
|
||||
# (Optional) Keep the same attribute for compatibility with splitting, etc.
|
||||
_fsdp_shard_conditions = [
|
||||
lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit(),
|
||||
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
|
||||
]
|
||||
_fsdp_shard_conditions = StepVideoConfig()._fsdp_shard_conditions
|
||||
_param_names_mapping = StepVideoConfig()._param_names_mapping
|
||||
_reverse_param_names_mapping = StepVideoConfig(
|
||||
)._reverse_param_names_mapping
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional, Tuple
|
||||
from dataclasses import field
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -12,6 +13,9 @@ from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class TextEncoder(nn.Module, ABC):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
|
||||
_stacked_params_mapping: List[Tuple[str, str,
|
||||
str]] = field(default_factory=list)
|
||||
_supported_attention_backends: Tuple[
|
||||
AttentionBackendEnum,
|
||||
...] = TextEncoderConfig()._supported_attention_backends
|
||||
@@ -19,6 +23,8 @@ class TextEncoder(nn.Module, ABC):
|
||||
def __init__(self, config: TextEncoderConfig) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self._fsdp_shard_conditions = config._fsdp_shard_conditions
|
||||
self._stacked_params_mapping = config.arch_config.stacked_params_mapping
|
||||
if not self.supported_attention_backends:
|
||||
raise ValueError(
|
||||
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
|
||||
|
||||
@@ -596,12 +596,7 @@ class CLIPVisionModel(ImageEncoder):
|
||||
# ref: https://github.com/vllm-project/vllm/pull/7186#discussion_r1734163986
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
]
|
||||
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
layer_count = len(self.vision_model.encoder.layers)
|
||||
@@ -620,7 +615,8 @@ class CLIPVisionModel(ImageEncoder):
|
||||
if layer_idx >= layer_count:
|
||||
continue
|
||||
|
||||
for (param_name, weight_name, shard_id) in stacked_params_mapping:
|
||||
for (param_name, weight_name,
|
||||
shard_id) in self.config.arch_config.stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
|
||||
@@ -369,14 +369,7 @@ class LlamaModel(TextEncoder):
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q_proj", "q"),
|
||||
(".qkv_proj", ".k_proj", "k"),
|
||||
(".qkv_proj", ".v_proj", "v"),
|
||||
(".gate_up_proj", ".gate_proj", 0),
|
||||
(".gate_up_proj", ".up_proj", 1),
|
||||
]
|
||||
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
@@ -406,7 +399,7 @@ class LlamaModel(TextEncoder):
|
||||
continue
|
||||
else:
|
||||
name = kv_scale_name
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
for param_name, weight_name, shard_id in self.config.arch_config.stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
|
||||
@@ -494,7 +494,7 @@ class T5Stack(nn.Module):
|
||||
attention_mask=attention_mask,
|
||||
attn_metadata=attn_metadata,
|
||||
)
|
||||
hidden_states = self.final_layer_norm.forward_native(hidden_states)
|
||||
hidden_states = self.final_layer_norm.forward(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -631,19 +631,13 @@ class UMT5EncoderModel(TextEncoder):
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q", "q"),
|
||||
(".qkv_proj", ".k", "k"),
|
||||
(".qkv_proj", ".v", "v"),
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
loaded = False
|
||||
if "decoder" in name or "lm_head" in name:
|
||||
continue
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
for param_name, weight_name, shard_id in self.config.arch_config.stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
|
||||
@@ -10,17 +10,20 @@ from copy import deepcopy
|
||||
from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from transformers import AutoImageProcessor, AutoTokenizer
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
from fastvideo.v1.configs.models import EncoderConfig
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
|
||||
from fastvideo.v1.models.loader.fsdp_load import maybe_load_fsdp_model
|
||||
from fastvideo.v1.models.loader.fsdp_load import (init_device_mesh,
|
||||
maybe_load_fsdp_model,
|
||||
shard_model)
|
||||
from fastvideo.v1.models.loader.utils import set_default_torch_dtype
|
||||
from fastvideo.v1.models.loader.weight_utils import (
|
||||
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
|
||||
@@ -163,16 +166,19 @@ class TextEncoderLoader(ComponentLoader):
|
||||
return hf_folder, hf_weights_files, use_safetensors
|
||||
|
||||
def _get_weights_iterator(
|
||||
self, source: "Source"
|
||||
self,
|
||||
source: "Source",
|
||||
to_cpu: bool = True
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Get an iterator for the model weights based on the load format."""
|
||||
hf_folder, hf_weights_files, use_safetensors = self._prepare_weights(
|
||||
source.model_or_path, source.fall_back_to_pt,
|
||||
source.allow_patterns_overrides)
|
||||
if use_safetensors:
|
||||
weights_iterator = safetensors_weights_iterator(hf_weights_files)
|
||||
weights_iterator = safetensors_weights_iterator(
|
||||
hf_weights_files, to_cpu)
|
||||
else:
|
||||
weights_iterator = pt_weights_iterator(hf_weights_files)
|
||||
weights_iterator = pt_weights_iterator(hf_weights_files, to_cpu)
|
||||
|
||||
if self.counter_before_loading_weights == 0.0:
|
||||
self.counter_before_loading_weights = time.perf_counter()
|
||||
@@ -181,10 +187,11 @@ class TextEncoderLoader(ComponentLoader):
|
||||
for (name, tensor) in weights_iterator)
|
||||
|
||||
def _get_all_weights(
|
||||
self,
|
||||
model_config: Any,
|
||||
model: nn.Module,
|
||||
model_path: str,
|
||||
self,
|
||||
model_config: Any,
|
||||
model: nn.Module,
|
||||
model_path: str,
|
||||
to_cpu: bool = True
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
primary_weights = TextEncoderLoader.Source(
|
||||
model_path,
|
||||
@@ -193,14 +200,14 @@ class TextEncoderLoader(ComponentLoader):
|
||||
allow_patterns_overrides=getattr(model, "allow_patterns_overrides",
|
||||
None),
|
||||
)
|
||||
yield from self._get_weights_iterator(primary_weights)
|
||||
yield from self._get_weights_iterator(primary_weights, to_cpu)
|
||||
|
||||
secondary_weights = cast(
|
||||
Iterable[TextEncoderLoader.Source],
|
||||
getattr(model, "secondary_weights", ()),
|
||||
)
|
||||
for source in secondary_weights:
|
||||
yield from self._get_weights_iterator(source)
|
||||
yield from self._get_weights_iterator(source, to_cpu)
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
fastvideo_args: FastVideoArgs):
|
||||
@@ -233,16 +240,22 @@ class TextEncoderLoader(ComponentLoader):
|
||||
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
|
||||
1]
|
||||
|
||||
target_device = get_torch_device()
|
||||
target_device = get_local_torch_device()
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, encoder_config, target_device,
|
||||
encoder_precision)
|
||||
fastvideo_args, encoder_precision)
|
||||
|
||||
def load_model(self,
|
||||
model_path: str,
|
||||
model_config: EncoderConfig,
|
||||
target_device: torch.device,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
dtype: str = "fp16"):
|
||||
use_cpu_offload = fastvideo_args.text_encoder_offload and len(
|
||||
getattr(model_config, "_fsdp_shard_conditions", [])) > 0
|
||||
|
||||
if fastvideo_args.text_encoder_offload:
|
||||
target_device = torch.device("cpu")
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
|
||||
with target_device:
|
||||
architectures = getattr(model_config, "architectures", [])
|
||||
@@ -251,12 +264,26 @@ class TextEncoderLoader(ComponentLoader):
|
||||
|
||||
weights_to_load = {name for name, _ in model.named_parameters()}
|
||||
loaded_weights = model.load_weights(
|
||||
self._get_all_weights(model_config, model, model_path))
|
||||
self._get_all_weights(model_config, model, model_path,
|
||||
use_cpu_offload))
|
||||
self.counter_after_loading_weights = time.perf_counter()
|
||||
logger.info(
|
||||
"Loading weights took %.2f seconds",
|
||||
self.counter_after_loading_weights -
|
||||
self.counter_before_loading_weights)
|
||||
|
||||
if use_cpu_offload:
|
||||
mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(1, dist.get_world_size()),
|
||||
mesh_dim_names=("offload", "replicate"),
|
||||
)
|
||||
shard_model(model,
|
||||
cpu_offload=True,
|
||||
reshard_after_forward=True,
|
||||
mesh=mesh["offload"],
|
||||
fsdp_shard_conditions=model._fsdp_shard_conditions,
|
||||
pin_cpu_memory=fastvideo_args.pin_cpu_memory)
|
||||
# We only enable strict check for non-quantized models
|
||||
# that have loaded weights tracking currently.
|
||||
# if loaded_weights is not None:
|
||||
@@ -290,10 +317,10 @@ class ImageEncoderLoader(TextEncoderLoader):
|
||||
encoder_config = fastvideo_args.pipeline_config.image_encoder_config
|
||||
encoder_config.update_model_arch(model_config)
|
||||
|
||||
target_device = get_torch_device()
|
||||
target_device = get_local_torch_device()
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(
|
||||
model_path, encoder_config, target_device,
|
||||
model_path, encoder_config, target_device, fastvideo_args,
|
||||
fastvideo_args.pipeline_config.image_encoder_precision)
|
||||
|
||||
|
||||
@@ -346,7 +373,7 @@ class VAELoader(ComponentLoader):
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]):
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
vae = vae_cls(vae_config).to(get_torch_device())
|
||||
vae = vae_cls(vae_config).to(get_local_torch_device())
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
@@ -405,7 +432,7 @@ class TransformerLoader(ComponentLoader):
|
||||
"hf_config": hf_config
|
||||
},
|
||||
weight_dir_list=safetensors_list,
|
||||
device=get_torch_device(),
|
||||
device=get_local_torch_device(),
|
||||
hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim,
|
||||
hsdp_shard_dim=fastvideo_args.hsdp_shard_dim,
|
||||
cpu_offload=fastvideo_args.use_cpu_offload,
|
||||
|
||||
@@ -69,6 +69,7 @@ def maybe_load_fsdp_model(
|
||||
fsdp_inference: bool = False,
|
||||
output_dtype: Optional[torch.dtype] = None,
|
||||
training_mode: bool = True,
|
||||
pin_cpu_memory: bool = True,
|
||||
) -> torch.nn.Module:
|
||||
"""
|
||||
Load the model with FSDP if is training, else load the model without FSDP.
|
||||
@@ -101,9 +102,12 @@ def maybe_load_fsdp_model(
|
||||
cpu_offload=cpu_offload,
|
||||
reshard_after_forward=True,
|
||||
mp_policy=mp_policy,
|
||||
mesh=device_mesh)
|
||||
mesh=device_mesh,
|
||||
fsdp_shard_conditions=model._fsdp_shard_conditions,
|
||||
pin_cpu_memory=pin_cpu_memory)
|
||||
|
||||
weight_iterator = safetensors_weights_iterator(weight_dir_list)
|
||||
weight_iterator = safetensors_weights_iterator(
|
||||
weight_dir_list, to_cpu=cpu_offload, async_broadcast=not cpu_offload)
|
||||
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
|
||||
load_model_from_full_model_state_dict(
|
||||
model,
|
||||
@@ -126,12 +130,13 @@ def maybe_load_fsdp_model(
|
||||
|
||||
def shard_model(
|
||||
model,
|
||||
*,
|
||||
cpu_offload: bool,
|
||||
reshard_after_forward: bool = True,
|
||||
mp_policy: Optional[MixedPrecisionPolicy] = None,
|
||||
dp_mesh: Optional[DeviceMesh] = None,
|
||||
mp_policy: Optional[MixedPrecisionPolicy] = MixedPrecisionPolicy(), # noqa
|
||||
mesh: Optional[DeviceMesh] = None,
|
||||
fsdp_shard_conditions: Optional[List[Callable[[str, nn.Module],
|
||||
bool]]] = None,
|
||||
pin_cpu_memory: bool = True,
|
||||
) -> None:
|
||||
"""
|
||||
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API.
|
||||
@@ -150,19 +155,28 @@ def shard_model(
|
||||
reshard_after_forward (bool): Whether to reshard parameters and buffers after
|
||||
the forward pass. Setting this to True corresponds to the FULL_SHARD sharding strategy
|
||||
from FSDP1, while setting it to False corresponds to the SHARD_GRAD_OP sharding strategy.
|
||||
dp_mesh (Optional[DeviceMesh]): Device mesh to use for FSDP sharding under multiple parallelism.
|
||||
mesh (Optional[DeviceMesh]): Device mesh to use for FSDP sharding under multiple parallelism.
|
||||
Default to None.
|
||||
fsdp_shard_conditions (Optional[List[Callable[[str, nn.Module], bool]]]): A list of functions to determine
|
||||
which modules to shard with FSDP.
|
||||
|
||||
Raises:
|
||||
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
|
||||
"""
|
||||
if fsdp_shard_conditions is None or len(fsdp_shard_conditions) == 0:
|
||||
logger.warning(
|
||||
"The FSDP shard condition list is empty or None. No modules will be sharded in %s",
|
||||
type(model).__name__)
|
||||
return
|
||||
|
||||
fsdp_kwargs = {
|
||||
"reshard_after_forward": reshard_after_forward,
|
||||
"mesh": mesh,
|
||||
"mp_policy": mp_policy,
|
||||
}
|
||||
if cpu_offload:
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(
|
||||
pin_memory=pin_cpu_memory)
|
||||
|
||||
# iterating in reverse to start with
|
||||
# lowest-level modules first
|
||||
@@ -172,7 +186,7 @@ def shard_model(
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([
|
||||
shard_condition(n, m)
|
||||
for shard_condition in model._fsdp_shard_conditions
|
||||
for shard_condition in fsdp_shard_conditions
|
||||
]):
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
num_layers_sharded += 1
|
||||
@@ -181,7 +195,6 @@ def shard_model(
|
||||
raise ValueError(
|
||||
"No layer modules were sharded. Please check if shard conditions are working as expected."
|
||||
)
|
||||
|
||||
# Finally shard the entire model to account for any stragglers
|
||||
fully_shard(model, **fsdp_kwargs)
|
||||
|
||||
@@ -224,6 +237,9 @@ def load_model_from_full_model_state_dict(
|
||||
to_merge_params: DefaultDict[str, Dict[Any, Any]] = defaultdict(dict)
|
||||
reverse_param_names_mapping = {}
|
||||
assert param_names_mapping is not None
|
||||
|
||||
# iterate over all the weights to sync broadcast before use
|
||||
full_sd_iterator = list(full_sd_iterator) # type: ignore
|
||||
for source_param_name, full_tensor in full_sd_iterator:
|
||||
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
|
||||
source_param_name)
|
||||
|
||||
@@ -11,9 +11,11 @@ from typing import Generator, List, Optional, Tuple, Union
|
||||
import filelock
|
||||
import huggingface_hub.constants
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from safetensors.torch import safe_open
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import get_node_group
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -118,36 +120,77 @@ _BAR_FORMAT = "{desc}: {percentage:3.0f}% Completed | {n_fmt}/{total_fmt} [{elap
|
||||
|
||||
|
||||
def safetensors_weights_iterator(
|
||||
hf_weights_files: List[str]
|
||||
hf_weights_files: List[str],
|
||||
to_cpu: bool = False,
|
||||
async_broadcast: bool = False
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Iterate over the weights in the model safetensor files."""
|
||||
enable_tqdm = not torch.distributed.is_initialized(
|
||||
) or torch.distributed.get_rank() == 0
|
||||
"""Iterate over the weights in the model safetensor files.
|
||||
Args:
|
||||
hf_weights_files: List of safetensor files to load.
|
||||
to_cpu: Whether to load the weights to CPU. If False, will load to the GPU device bound to the current process.
|
||||
async_broadcast: Whether to overlap loading from disk and broadcasting to other ranks. If True,
|
||||
must iterate over all the weights before use. Only use if to_cpu is False.
|
||||
"""
|
||||
local_rank = get_node_group().rank
|
||||
device = f"cuda:{local_rank}" if not to_cpu else "cpu"
|
||||
enable_tqdm = not torch.distributed.is_initialized() or get_node_group(
|
||||
).rank == 0
|
||||
assert not (async_broadcast
|
||||
and to_cpu), "Cannot broadcast weights when loading to CPU"
|
||||
|
||||
handles = []
|
||||
for st_file in tqdm(
|
||||
hf_weights_files,
|
||||
desc="Loading safetensors checkpoint shards",
|
||||
disable=not enable_tqdm,
|
||||
bar_format=_BAR_FORMAT,
|
||||
):
|
||||
with safe_open(st_file, framework="pt") as f:
|
||||
with safe_open(st_file, framework="pt", device=device) as f:
|
||||
for name in f.keys(): # noqa: SIM118
|
||||
param = f.get_tensor(name)
|
||||
if to_cpu:
|
||||
param = f.get_tensor(name)
|
||||
else:
|
||||
if local_rank == 0:
|
||||
param = f.get_tensor(name)
|
||||
else:
|
||||
shape = f.get_slice(name).get_shape()
|
||||
param = torch.empty(shape, device=device)
|
||||
# broadcast to local ranks
|
||||
# TODO(Wenxuan): scatter instead of broadcast
|
||||
if get_node_group().world_size > 1:
|
||||
group = get_node_group().device_group
|
||||
if async_broadcast:
|
||||
handle = dist.broadcast(param,
|
||||
src=dist.get_global_rank(
|
||||
group, 0),
|
||||
async_op=True)
|
||||
handles.append(handle)
|
||||
else:
|
||||
dist.broadcast(param,
|
||||
src=dist.get_global_rank(group, 0))
|
||||
yield name, param
|
||||
|
||||
if async_broadcast:
|
||||
for handle in handles:
|
||||
handle.wait()
|
||||
|
||||
|
||||
def pt_weights_iterator(
|
||||
hf_weights_files: List[str]
|
||||
hf_weights_files: List[str],
|
||||
to_cpu: bool = True # default to CPU for text encoder
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Iterate over the weights in the model bin/pt files."""
|
||||
enable_tqdm = not torch.distributed.is_initialized(
|
||||
) or torch.distributed.get_rank() == 0
|
||||
local_rank = get_node_group().rank
|
||||
device = f"cuda:{local_rank}" if not to_cpu else "cpu"
|
||||
enable_tqdm = not torch.distributed.is_initialized() or get_node_group(
|
||||
).rank == 0
|
||||
for bin_file in tqdm(
|
||||
hf_weights_files,
|
||||
desc="Loading pt checkpoint shards",
|
||||
disable=not enable_tqdm,
|
||||
bar_format=_BAR_FORMAT,
|
||||
):
|
||||
state = torch.load(bin_file, map_location="cpu", weights_only=True)
|
||||
state = torch.load(bin_file, map_location=device, weights_only=True)
|
||||
yield from state.items()
|
||||
del state
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ from tqdm import tqdm
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.dataset import ValidationDataset, getdataset
|
||||
from fastvideo.v1.dataset.preprocessing_datasets import PreprocessBatch
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
@@ -328,7 +328,8 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
# VAE
|
||||
with torch.autocast("cuda", dtype=torch.float32):
|
||||
latents = self.get_module("vae").encode(
|
||||
valid_data["pixel_values"].to(get_torch_device())).mean
|
||||
valid_data["pixel_values"].to(
|
||||
get_local_torch_device())).mean
|
||||
|
||||
# Get extra features if needed
|
||||
extra_features = self.get_extra_features(
|
||||
|
||||
@@ -13,7 +13,7 @@ import torch
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
@@ -82,8 +82,8 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
|
||||
# TODO(will): move these to cpu at some point
|
||||
self.get_module("image_encoder").to(get_torch_device())
|
||||
self.get_module("vae").to(get_torch_device())
|
||||
self.get_module("image_encoder").to(get_local_torch_device())
|
||||
self.get_module("vae").to(get_local_torch_device())
|
||||
|
||||
features = {}
|
||||
"""Get CLIP features from the first frame of each video."""
|
||||
@@ -107,7 +107,7 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
# Get CLIP features
|
||||
pixel_values = torch.cat(
|
||||
[img['pixel_values'] for img in processed_images],
|
||||
dim=0).to(get_torch_device())
|
||||
dim=0).to(get_local_torch_device())
|
||||
with torch.no_grad():
|
||||
image_inputs = {'pixel_values': pixel_values}
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
@@ -129,8 +129,8 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
height, width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=get_torch_device(),
|
||||
dtype=torch.float32)
|
||||
video_condition = video_condition.to(
|
||||
device=get_local_torch_device(), dtype=torch.float32)
|
||||
video_conditions.append(video_condition)
|
||||
|
||||
video_conditions = torch.cat(video_conditions, dim=0)
|
||||
|
||||
@@ -5,7 +5,7 @@ Decoding stage for diffusion pipelines.
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
@@ -61,7 +61,7 @@ class DecodingStage(PipelineStage):
|
||||
Returns:
|
||||
The batch with decoded outputs.
|
||||
"""
|
||||
self.vae = self.vae.to(get_torch_device())
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
latents = batch.latents
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
|
||||
@@ -12,8 +12,9 @@ from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.v1.attention import get_attn_backend
|
||||
from fastvideo.v1.configs.pipelines.base import STA_Mode
|
||||
from fastvideo.v1.distributed import (get_sp_parallel_rank, get_sp_world_size,
|
||||
get_torch_device, get_world_group)
|
||||
from fastvideo.v1.distributed import (get_local_torch_device,
|
||||
get_sp_parallel_rank, get_sp_world_size,
|
||||
get_world_group)
|
||||
from fastvideo.v1.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
@@ -192,7 +193,7 @@ class DenoisingStage(PipelineStage):
|
||||
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=get_torch_device(),
|
||||
device=get_local_torch_device(),
|
||||
).to(target_dtype) *
|
||||
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
|
||||
is not None else None)
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing import Optional
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
@@ -49,7 +49,7 @@ class EncodingStage(PipelineStage):
|
||||
Returns:
|
||||
The batch with encoded outputs.
|
||||
"""
|
||||
self.vae = self.vae.to(get_torch_device())
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
assert batch.height is not None
|
||||
assert batch.width is not None
|
||||
@@ -65,7 +65,8 @@ class EncodingStage(PipelineStage):
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=batch.height,
|
||||
width=batch.width).to(get_torch_device(), dtype=torch.float32)
|
||||
width=batch.width).to(get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
image = image.unsqueeze(2)
|
||||
else:
|
||||
@@ -78,7 +79,7 @@ class EncodingStage(PipelineStage):
|
||||
batch.num_frames - 1, batch.height, batch.width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=get_torch_device(),
|
||||
video_condition = video_condition.to(device=get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
# Setup VAE precision
|
||||
@@ -102,6 +103,7 @@ class EncodingStage(PipelineStage):
|
||||
generator = batch.generator
|
||||
if generator is None:
|
||||
raise ValueError("Generator must be provided")
|
||||
# latent_condition = self.retrieve_latents(encoder_output, generator, sample_mode="argmax")
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator)
|
||||
|
||||
# Apply shifting if needed
|
||||
|
||||
@@ -7,7 +7,7 @@ This module contains implementations of image encoding stages for diffusion pipe
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -55,12 +55,12 @@ class ImageEncodingStage(PipelineStage):
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder = self.image_encoder.to(get_torch_device())
|
||||
self.image_encoder = self.image_encoder.to(get_local_torch_device())
|
||||
|
||||
image = batch.pil_image
|
||||
|
||||
image_inputs = self.image_processor(
|
||||
images=image, return_tensors="pt").to(get_torch_device())
|
||||
images=image, return_tensors="pt").to(get_local_torch_device())
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = self.image_encoder(**image_inputs)
|
||||
image_embeds = outputs.last_hidden_state
|
||||
|
||||
@@ -5,7 +5,7 @@ Latent preparation stage for diffusion pipelines.
|
||||
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -62,7 +62,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
# Get required parameters
|
||||
dtype = batch.prompt_embeds[0].dtype
|
||||
device = get_torch_device()
|
||||
device = get_local_torch_device()
|
||||
generator = batch.generator
|
||||
latents = batch.latents
|
||||
num_frames = latent_num_frames if latent_num_frames is not None else batch.num_frames
|
||||
|
||||
@@ -5,9 +5,7 @@ Prompt encoding stages for diffusion pipelines.
|
||||
This module contains implementations of prompt encoding stages for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -62,8 +60,6 @@ class TextEncodingStage(PipelineStage):
|
||||
fastvideo_args.pipeline_config.text_encoder_configs,
|
||||
fastvideo_args.pipeline_config.preprocess_text_funcs,
|
||||
fastvideo_args.pipeline_config.postprocess_text_funcs):
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
text_encoder = text_encoder.to(get_torch_device())
|
||||
|
||||
assert isinstance(batch.prompt, (str, list))
|
||||
if isinstance(batch.prompt, str):
|
||||
@@ -71,8 +67,9 @@ class TextEncodingStage(PipelineStage):
|
||||
texts = []
|
||||
for prompt_str in batch.prompt:
|
||||
texts.append(preprocess_func(prompt_str))
|
||||
text_inputs = tokenizer(
|
||||
texts, **encoder_config.tokenizer_kwargs).to(get_torch_device())
|
||||
text_inputs = tokenizer(texts,
|
||||
**encoder_config.tokenizer_kwargs).to(
|
||||
get_local_torch_device())
|
||||
input_ids = text_inputs["input_ids"]
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
@@ -91,8 +88,8 @@ class TextEncodingStage(PipelineStage):
|
||||
assert isinstance(batch.negative_prompt, str)
|
||||
negative_text = preprocess_func(batch.negative_prompt)
|
||||
negative_text_inputs = tokenizer(
|
||||
negative_text,
|
||||
**encoder_config.tokenizer_kwargs).to(get_torch_device())
|
||||
negative_text, **encoder_config.tokenizer_kwargs).to(
|
||||
get_local_torch_device())
|
||||
negative_input_ids = negative_text_inputs["input_ids"]
|
||||
negative_attention_mask = negative_text_inputs["attention_mask"]
|
||||
with set_forward_context(current_timestep=0,
|
||||
@@ -110,9 +107,8 @@ class TextEncodingStage(PipelineStage):
|
||||
batch.negative_attention_mask.append(
|
||||
negative_attention_mask)
|
||||
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
if fastvideo_args.text_encoder_offload:
|
||||
text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ This module contains implementations of timestep preparation stages for diffusio
|
||||
|
||||
import inspect
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -45,7 +45,7 @@ class TimestepPreparationStage(PipelineStage):
|
||||
The batch with prepared timesteps.
|
||||
"""
|
||||
scheduler = self.scheduler
|
||||
device = get_torch_device()
|
||||
device = get_local_torch_device()
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
timesteps = batch.timesteps
|
||||
sigmas = batch.sigmas
|
||||
|
||||
@@ -14,7 +14,7 @@ from typing import Any, Dict
|
||||
import torch
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.encoders.bert import HunyuanClip # type: ignore
|
||||
@@ -78,7 +78,7 @@ class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
target_device = get_torch_device()
|
||||
target_device = get_local_torch_device()
|
||||
llm_dir = os.path.join(self.model_path, "step_llm")
|
||||
clip_dir = os.path.join(self.model_path, "hunyuan_clip")
|
||||
text_enc = self.build_llm(llm_dir, target_device)
|
||||
|
||||
@@ -6,7 +6,7 @@ import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from transformers import AutoConfig
|
||||
|
||||
import gc
|
||||
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
# from fastvideo.v1.models.hunyuan.text_encoder import load_text_encoder, load_tokenizer
|
||||
@@ -16,6 +16,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.configs.models.encoders import CLIPTextConfig
|
||||
from torch.distributed.tensor import DTensor
|
||||
from torch.testing import assert_close
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -66,7 +68,6 @@ def test_clip_encoder():
|
||||
# Load the HuggingFace implementation directly
|
||||
# model2 = CLIPTextModel(hf_config)
|
||||
# model2 = model2.to(torch.float16)
|
||||
model2 = model2.to(device)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights between the two models
|
||||
@@ -78,19 +79,20 @@ def test_clip_encoder():
|
||||
logger.info("Model1 has %d parameters", len(params1))
|
||||
logger.info("Model2 has %d parameters", len(params2))
|
||||
|
||||
# Compare a few key parameters
|
||||
|
||||
# weight_diffs = []
|
||||
# for (name1, param1), (name2, param2) in zip(
|
||||
# sorted(params1.items()), sorted(params2.items())
|
||||
# ):
|
||||
# # if len(weight_diffs) < 5: # Just check a few parameters
|
||||
# max_diff = torch.max(torch.abs(param1 - param2)).item()
|
||||
# mean_diff = torch.mean(torch.abs(param1 - param2)).item()
|
||||
# weight_diffs.append((name1, name2, max_diff, mean_diff))
|
||||
# logger.info(f"Parameter: {name1} vs {name2}")
|
||||
# logger.info(f" Max diff: {max_diff}, Mean diff: {mean_diff}")
|
||||
|
||||
for name1, param1 in sorted(params1.items()):
|
||||
name2 = name1
|
||||
skip = False
|
||||
for param_name, weight_name, shard_id in model2.config.arch_config.stacked_params_mapping:
|
||||
if weight_name not in name1:
|
||||
skip = True
|
||||
# stacked params are more troublesome
|
||||
if skip:
|
||||
continue
|
||||
param2 = params2[name2]
|
||||
param2 = param2.to_local().to(device) if isinstance(param2, DTensor) else param2.to(device)
|
||||
assert_close(param1, param2, atol=1e-4, rtol=1e-4)
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
# Load tokenizer
|
||||
tokenizer, _ = load_tokenizer(tokenizer_type="clipL",
|
||||
tokenizer_path=args.model_path,
|
||||
|
||||
@@ -5,7 +5,7 @@ import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from transformers import AutoConfig
|
||||
|
||||
import gc
|
||||
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
@@ -15,7 +15,8 @@ from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.configs.models.encoders import LlamaConfig
|
||||
|
||||
from torch.distributed.tensor import DTensor
|
||||
from torch.testing import assert_close
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
@@ -62,7 +63,6 @@ def test_llama_encoder():
|
||||
|
||||
# Convert to float16 and move to device
|
||||
# model2 = model2.to(torch.float16)
|
||||
model2 = model2.to(device)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights between the two models
|
||||
@@ -77,34 +77,28 @@ def test_llama_encoder():
|
||||
# Compare a few key parameters
|
||||
weight_diffs = []
|
||||
# check if embed_tokens are the same
|
||||
print(model1.embed_tokens.weight.shape, model2.embed_tokens.weight.shape)
|
||||
device = model1.embed_tokens.weight.device
|
||||
assert torch.allclose(model1.embed_tokens.weight,
|
||||
model2.embed_tokens.weight)
|
||||
model2.embed_tokens.weight.to_local().to(device) if isinstance(model2.embed_tokens.weight, DTensor) else model2.embed_tokens.weight.to(device))
|
||||
weights = [
|
||||
"layers.{}.input_layernorm.weight",
|
||||
"layers.{}.post_attention_layernorm.weight"
|
||||
]
|
||||
# for (name1, param1), (name2, param2) in zip(
|
||||
# sorted(params1.items()), sorted(params2.items())
|
||||
# ):
|
||||
for layer_idx in range(hf_config.num_hidden_layers):
|
||||
for w in weights:
|
||||
name1 = w.format(layer_idx)
|
||||
name2 = w.format(layer_idx)
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
# print(type(p2))
|
||||
if "gate_up" in name2:
|
||||
# print("skipping gate_up")
|
||||
continue
|
||||
try:
|
||||
# logger.info(f"Parameter: {name1} vs {name2}")
|
||||
max_diff = torch.max(torch.abs(p1 - p2)).item()
|
||||
mean_diff = torch.mean(torch.abs(p1 - p2)).item()
|
||||
weight_diffs.append((name1, name2, max_diff, mean_diff))
|
||||
# logger.info(f" Max diff: {max_diff}, Mean diff: {mean_diff}")
|
||||
except Exception as e:
|
||||
logger.info("Error comparing %s and %s: %s", name1, name2, e)
|
||||
|
||||
for name1, param1 in sorted(params1.items()):
|
||||
name2 = name1
|
||||
skip = False
|
||||
for param_name, weight_name, shard_id in model2.config.arch_config.stacked_params_mapping:
|
||||
if weight_name not in name1:
|
||||
skip = True
|
||||
# stacked params are more troublesome
|
||||
if skip:
|
||||
continue
|
||||
param2 = params2[name2]
|
||||
param2 = param2.to_local().to(device) if isinstance(param2, DTensor) else param2.to(device)
|
||||
assert_close(param1, param2, atol=1e-4, rtol=1e-4)
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
tokenizer, _ = load_tokenizer(tokenizer_type="llm",
|
||||
tokenizer_path=TOKENIZER_PATH,
|
||||
|
||||
@@ -4,6 +4,8 @@ import os
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from torch.distributed.tensor import DTensor
|
||||
from torch.testing import assert_close
|
||||
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel
|
||||
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
@@ -41,13 +43,13 @@ def test_t5_encoder():
|
||||
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
|
||||
|
||||
|
||||
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH, pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),), text_encoder_precisions=(precision_str,)))
|
||||
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH,
|
||||
pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),),
|
||||
text_encoder_precisions=(precision_str,)),
|
||||
pin_cpu_memory=False)
|
||||
loader = TextEncoderLoader()
|
||||
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
|
||||
|
||||
# Convert to float16 and move to device
|
||||
# model2 = model2.to(precision)
|
||||
model2 = model2.to(device)
|
||||
model2 = model2.to(precision)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights between the two models
|
||||
@@ -64,23 +66,17 @@ def test_t5_encoder():
|
||||
weights = ["encoder.block.{}.layer.0.layer_norm.weight", "encoder.block.{}.layer.0.SelfAttention.relative_attention_bias.weight", \
|
||||
"encoder.block.{}.layer.0.SelfAttention.o.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_0.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_1.weight",\
|
||||
"encoder.block.{}.layer.1.DenseReluDense.wo.weight", \
|
||||
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight", "shared.weight"]
|
||||
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight"]
|
||||
|
||||
for idx in range(hf_config.num_hidden_layers):
|
||||
for w in weights:
|
||||
name1 = w.format(idx)
|
||||
name2 = w.format(idx)
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
assert p1.dtype == p2.dtype
|
||||
try:
|
||||
logger.info("Parameter: %s vs %s", name1, name2)
|
||||
max_diff = torch.max(torch.abs(p1 - p2)).item()
|
||||
mean_diff = torch.mean(torch.abs(p1 - p2)).item()
|
||||
weight_diffs.append((name1, name2, max_diff, mean_diff))
|
||||
logger.info(" Max diff: %s, Mean diff: %s", max_diff,
|
||||
mean_diff)
|
||||
except Exception as e:
|
||||
logger.info("Error comparing %s and %s: %s", name1, name2, e)
|
||||
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(p1)
|
||||
assert_close(p1, p2, atol=1e-4, rtol=1e-4)
|
||||
|
||||
|
||||
# Test with some sample prompts
|
||||
prompts = [
|
||||
|
||||
@@ -5,7 +5,7 @@ from pathlib import Path
|
||||
import pytest
|
||||
|
||||
NUM_NODES = "1"
|
||||
NUM_GPUS_PER_NODE = "2"
|
||||
NUM_GPUS_PER_NODE = "1"
|
||||
|
||||
# Set environment variables
|
||||
os.environ["FASTVIDEO_ATTENTION_CONFIG"] = "assets/mask_strategy_wan.json"
|
||||
@@ -17,9 +17,9 @@ def test_inference():
|
||||
cmd = [
|
||||
"fastvideo", "generate",
|
||||
"--model-path", "Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
"--sp-size", "2",
|
||||
"--tp-size", "2",
|
||||
"--num-gpus", "2",
|
||||
"--sp-size", "1",
|
||||
"--tp-size", "1",
|
||||
"--num-gpus", "1",
|
||||
"--height", "768",
|
||||
"--width", "1280",
|
||||
"--num-frames", "69",
|
||||
|
||||
@@ -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:1", 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:1", 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 +1 @@
|
||||
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":0.50390625,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.19960195198655128,"_runtime":107.325113071}
|
||||
{"grad_norm":0.478515625,"_runtime":95.727033597,"_wandb":{"runtime":95},"_step":5,"validation_videos_50_steps":{"videos":[{"_type":"video-file","sha256":"42a1c311521a9d460db788713be1cbf2db767494e02619b43be5bf3eed8381d8","size":158632,"path":"media/videos/validation_videos_50_steps_0_42a1c311521a9d460db7.mp4"},{"path":"media/videos/validation_videos_50_steps_0_818505095b4b5e8b7f51.mp4","_type":"video-file","sha256":"818505095b4b5e8b7f511012d45f04d151ce3344bc058fc0f3225a414a851e4a","size":147825},{"sha256":"fc334ba9ed5e66c8527ee3b408e3be2d76167fef03588bf2840f4a0792f2fe34","size":136933,"path":"media/videos/validation_videos_50_steps_0_fc334ba9ed5e66c8527e.mp4","_type":"video-file"},{"size":201797,"path":"media/videos/validation_videos_50_steps_0_ccd98f6f907635d266a7.mp4","_type":"video-file","sha256":"ccd98f6f907635d266a74783688e7ecf1dac752d79d72d69eab9ef0e3f7413eb"},{"_type":"video-file","sha256":"ca79f40a0aed38f676f12779b349ce40e9e3fb7f36c578f49a20854c70508fb4","size":147114,"path":"media/videos/validation_videos_50_steps_0_ca79f40a0aed38f676f1.mp4"},{"size":175104,"path":"media/videos/validation_videos_50_steps_0_32c9b33ff920c17e5881.mp4","_type":"video-file","sha256":"32c9b33ff920c17e588133d7a27aa400ff3dc529b01ed4f16ac4d6bb2afa0f00"},{"sha256":"2cf520bfb93401c914e93c87ef791c2f12a4e043b95dfdc98115c930e11dfe67","size":139655,"path":"media/videos/validation_videos_50_steps_0_2cf520bfb93401c914e9.mp4","_type":"video-file"},{"_type":"video-file","sha256":"1d73aba17ce582c7aef4af4d64079e3e9d3df205634eff453446bdaf2340b214","size":149028,"path":"media/videos/validation_videos_50_steps_0_1d73aba17ce582c7aef4.mp4"}],"captions":false,"_type":"videos","count":8},"train_loss":0.08922439813613892,"_timestamp":1.750202051751466e+09,"avg_step_time":0.7536672964692116,"step_time":0.4742048177868128,"learning_rate":1e-05,"vsa_sparsity":0.05}
|
||||
@@ -15,7 +15,7 @@ wandb_name = "test_training_loss_VSA"
|
||||
reference_wandb_summary_file = "fastvideo/v1/tests/training/VSA/reference_wandb_summary_VSA.json"
|
||||
|
||||
NUM_NODES = "1"
|
||||
NUM_GPUS_PER_NODE = "2"
|
||||
NUM_GPUS_PER_NODE = "1"
|
||||
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
|
||||
@@ -35,14 +35,14 @@ def run_worker():
|
||||
"--validation_preprocessed_path", "data/mini_dataset_i2v_VSA/validation_parquet_dataset",
|
||||
"--train_batch_size", "1",
|
||||
"--num_latent_t", "4",
|
||||
"--num_gpus", "2",
|
||||
"--sp_size", "2",
|
||||
"--tp_size", "2",
|
||||
"--num_gpus", "1",
|
||||
"--sp_size", "1",
|
||||
"--tp_size", "1",
|
||||
"--hsdp_replicate_dim", "1",
|
||||
"--hsdp_shard_dim", "2",
|
||||
"--hsdp_shard_dim", "1",
|
||||
"--train_sp_batch_size", "1",
|
||||
"--dataloader_num_workers", "4",
|
||||
"--gradient_accumulation_steps", "2",
|
||||
"--gradient_accumulation_steps", "1",
|
||||
"--max_train_steps", "5",
|
||||
"--learning_rate", "1e-5",
|
||||
"--mixed_precision", "bf16",
|
||||
@@ -110,7 +110,7 @@ def test_distributed_training():
|
||||
fields_and_thresholds = {
|
||||
'avg_step_time': 1.0,
|
||||
'grad_norm': 0.1,
|
||||
'step_time': 1.0,
|
||||
'step_time': 0.5,
|
||||
'train_loss': 0.001
|
||||
}
|
||||
|
||||
|
||||
@@ -121,9 +121,9 @@ def test_distributed_training():
|
||||
wandb_summary = json.load(open(summary_file))
|
||||
|
||||
fields_and_thresholds = {
|
||||
'avg_step_time': 6.0,
|
||||
'avg_step_time': 1.0,
|
||||
'grad_norm': 0.3,
|
||||
'step_time': 6.0,
|
||||
'step_time': 0.5,
|
||||
'train_loss': 0.0025
|
||||
}
|
||||
|
||||
|
||||
@@ -80,7 +80,10 @@ def test_hunyuanvideo_distributed():
|
||||
|
||||
# Initialize with identical weights
|
||||
model = initialize_identical_weights(model, seed=42)
|
||||
shard_model(model, cpu_offload=False, reshard_after_forward=True)
|
||||
shard_model(model, cpu_offload=True,
|
||||
reshard_after_forward=True,
|
||||
fsdp_shard_conditions=model._fsdp_shard_conditions
|
||||
)
|
||||
for n, p in chain(model.named_parameters(), model.named_buffers()):
|
||||
if p.is_meta:
|
||||
raise RuntimeError(
|
||||
|
||||
@@ -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)
|
||||
@@ -24,15 +24,14 @@ from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.dataset import build_parquet_map_style_dataloader
|
||||
from fastvideo.v1.dataset.dataloader.schema import (
|
||||
pyarrow_schema_t2v, pyarrow_schema_t2v_validation)
|
||||
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
|
||||
get_torch_device, get_world_group)
|
||||
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory,
|
||||
get_local_torch_device, get_sp_group,
|
||||
get_world_group)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
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,
|
||||
@@ -67,8 +66,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing training pipeline...")
|
||||
self.device = get_local_torch_device()
|
||||
self.training_args = training_args
|
||||
self.device = get_torch_device()
|
||||
world_group = get_world_group()
|
||||
self.world_size = world_group.world_size
|
||||
self.global_rank = world_group.rank
|
||||
@@ -84,11 +83,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()
|
||||
@@ -183,12 +177,12 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
encoder_attention_mask = batch['text_attention_mask']
|
||||
infos = batch['info_list']
|
||||
|
||||
training_batch.latents = latents.to(get_torch_device(),
|
||||
training_batch.latents = latents.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
training_batch.encoder_hidden_states = encoder_hidden_states.to(
|
||||
get_torch_device(), dtype=torch.bfloat16)
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
get_torch_device(), dtype=torch.bfloat16)
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch
|
||||
@@ -252,8 +246,9 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
current_vsa_sparsity = training_batch.current_vsa_sparsity
|
||||
|
||||
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
|
||||
|
||||
dit_seq_shape = [
|
||||
latents.shape[2] * self.sp_world_size // patch_size[0],
|
||||
latents.shape[2] // patch_size[0],
|
||||
latents.shape[3] // patch_size[1],
|
||||
latents.shape[4] // patch_size[2]
|
||||
]
|
||||
@@ -280,7 +275,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
"encoder_hidden_states":
|
||||
training_batch.encoder_hidden_states,
|
||||
"timestep":
|
||||
training_batch.timesteps.to(get_torch_device(),
|
||||
training_batch.timesteps.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16),
|
||||
"encoder_attention_mask":
|
||||
training_batch.encoder_attention_mask,
|
||||
@@ -316,18 +311,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,14 +548,13 @@ 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']
|
||||
|
||||
prompt_embeds = prompt_embeds.to(get_torch_device())
|
||||
prompt_attention_mask = prompt_attention_mask.to(get_torch_device())
|
||||
prompt_embeds = prompt_embeds.to(get_local_torch_device())
|
||||
prompt_attention_mask = prompt_attention_mask.to(
|
||||
get_local_torch_device())
|
||||
|
||||
# Calculate sizes
|
||||
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
|
||||
@@ -639,7 +632,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 +640,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
|
||||
|
||||
@@ -8,7 +8,7 @@ import torch
|
||||
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_local_torch_device
|
||||
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 (
|
||||
@@ -85,15 +85,17 @@ class WanI2VTrainingPipeline(TrainingPipeline):
|
||||
pil_image = batch['pil_image']
|
||||
infos = batch['info_list']
|
||||
|
||||
training_batch.latents = latents.to(get_torch_device(),
|
||||
training_batch.latents = latents.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
training_batch.encoder_hidden_states = encoder_hidden_states.to(
|
||||
get_torch_device(), dtype=torch.bfloat16)
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
get_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.preprocessed_image = pil_image.to(get_torch_device())
|
||||
training_batch.image_embeds = clip_features.to(get_torch_device())
|
||||
training_batch.image_latents = image_latents.to(get_torch_device())
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.preprocessed_image = pil_image.to(
|
||||
get_local_torch_device())
|
||||
training_batch.image_embeds = clip_features.to(get_local_torch_device())
|
||||
training_batch.image_latents = image_latents.to(
|
||||
get_local_torch_device())
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch
|
||||
@@ -112,8 +114,8 @@ class WanI2VTrainingPipeline(TrainingPipeline):
|
||||
training_batch = super()._prepare_dit_inputs(training_batch)
|
||||
|
||||
assert isinstance(training_batch.image_latents, torch.Tensor)
|
||||
image_latents = training_batch.image_latents.to(get_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
image_latents = training_batch.image_latents.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
|
||||
training_batch.noisy_model_input = torch.cat(
|
||||
[training_batch.noisy_model_input, image_latents], dim=1)
|
||||
@@ -132,7 +134,8 @@ class WanI2VTrainingPipeline(TrainingPipeline):
|
||||
# Image Embeds for conditioning
|
||||
image_embeds = training_batch.image_embeds
|
||||
assert torch.isnan(image_embeds).sum() == 0
|
||||
image_embeds = image_embeds.to(get_torch_device(), dtype=torch.bfloat16)
|
||||
image_embeds = image_embeds.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
encoder_hidden_states_image = image_embeds
|
||||
|
||||
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
|
||||
@@ -142,7 +145,7 @@ class WanI2VTrainingPipeline(TrainingPipeline):
|
||||
"encoder_hidden_states":
|
||||
training_batch.encoder_hidden_states,
|
||||
"timestep":
|
||||
training_batch.timesteps.to(get_torch_device(),
|
||||
training_batch.timesteps.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16),
|
||||
"encoder_attention_mask":
|
||||
training_batch.encoder_attention_mask,
|
||||
@@ -166,9 +169,9 @@ class WanI2VTrainingPipeline(TrainingPipeline):
|
||||
infos = validation_batch['info_list']
|
||||
prompt = infos[0]['prompt']
|
||||
|
||||
prompt_embeds = embeddings.to(get_torch_device())
|
||||
prompt_attention_mask = masks.to(get_torch_device())
|
||||
clip_features = clip_features.to(get_torch_device())
|
||||
prompt_embeds = embeddings.to(get_local_torch_device())
|
||||
prompt_attention_mask = masks.to(get_local_torch_device())
|
||||
clip_features = clip_features.to(get_local_torch_device())
|
||||
|
||||
# Calculate sizes
|
||||
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
|
||||
|
||||
+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