Compare commits

...
43 changed files with 278 additions and 423 deletions
+5 -5
View File
@@ -122,17 +122,17 @@ jobs:
# Actual tests
encoder-test:
- 'fastvideo/v1/models/encoders/**'
- 'fastvideo/v1/models/loader/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/encoders/**'
- *common-paths
vae-test:
- 'fastvideo/v1/models/vaes/**'
- 'fastvideo/v1/models/loader/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/vaes/**'
- *common-paths
transformer-test:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loader/**'
- 'fastvideo/v1/models/loaders/**'
- '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: 1
gpu_count: 2
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: 1
gpu_count: 2
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
+1 -1
View File
@@ -10,7 +10,7 @@ def main():
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
# FastVideo will automatically handle distributed setup
# if num_gpus > 1, FastVideo will automatically handle distributed setup
num_gpus=2,
use_fsdp_inference=True,
use_cpu_offload=False
+2 -4
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field, fields
from typing import Any, Dict, List, Tuple
from typing import Any, Dict
from fastvideo.v1.logger import init_logger
@@ -12,9 +12,7 @@ logger = init_logger(__name__)
# 3. Any field in ArchConfig is fixed upon initialization, and should be hidden away from users
@dataclass
class ArchConfig:
stacked_params_mapping: List[Tuple[str, str, str]] = field(
default_factory=list
) # mapping from huggingface weight names to custom names
pass
@dataclass
@@ -5,11 +5,13 @@ 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:
[lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit()])
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
_param_names_mapping: dict = field(
default_factory=lambda: {
+1 -4
View File
@@ -32,11 +32,8 @@ 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 -18
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
ImageEncoderConfig,
@@ -8,14 +8,6 @@ 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
@@ -35,15 +27,6 @@ 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 -25
View File
@@ -1,23 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
from typing import Optional
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
@@ -44,18 +32,6 @@ 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 -23
View File
@@ -1,23 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
from typing import Optional
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
@@ -41,16 +29,6 @@ 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_local_torch_device,
cleanup_dist_env_and_memory, get_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_local_torch_device())
embeddings = embeddings.to(get_local_torch_device())
latents = latents.to(get_torch_device())
embeddings = embeddings.to(get_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_local_torch_device,
cleanup_dist_env_and_memory, get_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_local_torch_device())
embeddings = embeddings.to(get_local_torch_device())
latents = latents.to(get_torch_device())
embeddings = embeddings.to(get_torch_device())
# Calculate actual batch size
batch_size = latents.size(0)
+5 -5
View File
@@ -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_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,
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,
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_local_torch_device",
"get_torch_device",
]
+6 -32
View File
@@ -36,7 +36,6 @@ 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
@@ -693,7 +692,6 @@ class GroupCoordinator:
_WORLD: Optional[GroupCoordinator] = None
_NODE: Optional[GroupCoordinator] = None
def get_world_group() -> GroupCoordinator:
@@ -701,11 +699,6 @@ 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(
@@ -717,18 +710,6 @@ 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,
@@ -801,8 +782,6 @@ 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
@@ -925,7 +904,7 @@ def get_dp_rank() -> int:
return get_dp_group().rank_in_group
def get_local_torch_device() -> torch.device:
def get_torch_device() -> torch.device:
"""Return the torch device for the current rank."""
return torch.device(f"cuda:{envs.LOCAL_RANK}")
@@ -1042,22 +1021,17 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
"torch._C._host_emptyCache() only available in Pytorch >=2.5")
def same_node_ranks(pg: Union[ProcessGroup, StatelessProcessGroup],
source_rank: int = 0) -> List[int]:
def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
source_rank: int = 0) -> List[bool]:
"""
This is a collective operation that returns ranks that are in the same node
This is a collective operation that returns if each rank is 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, (
"same_node_ranks should be tested with a non-NCCL group.")
"in_the_same_node_as 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)
@@ -1129,7 +1103,7 @@ def same_node_ranks(pg: Union[ProcessGroup, StatelessProcessGroup],
rank_data = pg.broadcast_obj(is_in_the_same_node, src=i)
aggregated_data += rank_data
return [i for i, x in enumerate(aggregated_data.tolist()) if x == 1]
return [x == 1 for x in aggregated_data.tolist()]
def initialize_tensor_parallel_group(
+3 -17
View File
@@ -58,10 +58,8 @@ class FastVideoArgs:
output_type: str = "pil"
use_cpu_offload: bool = True # For DiT
use_cpu_offload: bool = True
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
@@ -210,7 +208,7 @@ class FastVideoArgs:
"--use-cpu-offload",
action=StoreBoolean,
help=
"Use CPU offload for DiT inference. Enable if run out of memory with FSDP.",
"Use CPU offload for model inference. Enable if run out of memory with FSDP.",
)
parser.add_argument(
"--use-fsdp-inference",
@@ -218,19 +216,7 @@ 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,
+7 -7
View File
@@ -6,7 +6,6 @@ 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
@@ -39,6 +38,12 @@ 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,
@@ -71,12 +76,7 @@ class RMSNorm(CustomOp):
x = x * torch.rsqrt(variance + self.variance_epsilon)
x = x.to(orig_dtype)
if self.has_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
x = x * self.weight
if residual is None:
return x
else:
+4 -1
View File
@@ -455,7 +455,10 @@ class StepVideoTransformerBlock(nn.Module):
class StepVideoModel(BaseDiT):
# (Optional) Keep the same attribute for compatibility with splitting, etc.
_fsdp_shard_conditions = StepVideoConfig()._fsdp_shard_conditions
_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.
]
_param_names_mapping = StepVideoConfig()._param_names_mapping
_reverse_param_names_mapping = StepVideoConfig(
)._reverse_param_names_mapping
+1 -7
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from dataclasses import field
from typing import List, Optional, Tuple
from typing import Optional, Tuple
import torch
from torch import nn
@@ -13,9 +12,6 @@ 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
@@ -23,8 +19,6 @@ 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"
+7 -3
View File
@@ -596,7 +596,12 @@ 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)
@@ -615,8 +620,7 @@ class CLIPVisionModel(ImageEncoder):
if layer_idx >= layer_count:
continue
for (param_name, weight_name,
shard_id) in self.config.arch_config.stacked_params_mapping:
for (param_name, weight_name, shard_id) in stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
+9 -2
View File
@@ -369,7 +369,14 @@ 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:
@@ -399,7 +406,7 @@ class LlamaModel(TextEncoder):
continue
else:
name = kv_scale_name
for param_name, weight_name, shard_id in self.config.arch_config.stacked_params_mapping:
for param_name, weight_name, shard_id in stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
+8 -2
View File
@@ -494,7 +494,7 @@ class T5Stack(nn.Module):
attention_mask=attention_mask,
attn_metadata=attn_metadata,
)
hidden_states = self.final_layer_norm.forward(hidden_states)
hidden_states = self.final_layer_norm.forward_native(hidden_states)
return hidden_states
@@ -631,13 +631,19 @@ 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 self.config.arch_config.stacked_params_mapping:
for param_name, weight_name, shard_id in stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
+18 -45
View File
@@ -10,20 +10,17 @@ 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_local_torch_device
from fastvideo.v1.distributed import get_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 (init_device_mesh,
maybe_load_fsdp_model,
shard_model)
from fastvideo.v1.models.loader.fsdp_load import maybe_load_fsdp_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,
@@ -166,19 +163,16 @@ class TextEncoderLoader(ComponentLoader):
return hf_folder, hf_weights_files, use_safetensors
def _get_weights_iterator(
self,
source: "Source",
to_cpu: bool = True
self, source: "Source"
) -> 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, to_cpu)
weights_iterator = safetensors_weights_iterator(hf_weights_files)
else:
weights_iterator = pt_weights_iterator(hf_weights_files, to_cpu)
weights_iterator = pt_weights_iterator(hf_weights_files)
if self.counter_before_loading_weights == 0.0:
self.counter_before_loading_weights = time.perf_counter()
@@ -187,11 +181,10 @@ class TextEncoderLoader(ComponentLoader):
for (name, tensor) in weights_iterator)
def _get_all_weights(
self,
model_config: Any,
model: nn.Module,
model_path: str,
to_cpu: bool = True
self,
model_config: Any,
model: nn.Module,
model_path: str,
) -> Generator[Tuple[str, torch.Tensor], None, None]:
primary_weights = TextEncoderLoader.Source(
model_path,
@@ -200,14 +193,14 @@ class TextEncoderLoader(ComponentLoader):
allow_patterns_overrides=getattr(model, "allow_patterns_overrides",
None),
)
yield from self._get_weights_iterator(primary_weights, to_cpu)
yield from self._get_weights_iterator(primary_weights)
secondary_weights = cast(
Iterable[TextEncoderLoader.Source],
getattr(model, "secondary_weights", ()),
)
for source in secondary_weights:
yield from self._get_weights_iterator(source, to_cpu)
yield from self._get_weights_iterator(source)
def load(self, model_path: str, architecture: str,
fastvideo_args: FastVideoArgs):
@@ -240,22 +233,16 @@ class TextEncoderLoader(ComponentLoader):
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
1]
target_device = get_local_torch_device()
target_device = get_torch_device()
# TODO(will): add support for other dtypes
return self.load_model(model_path, encoder_config, target_device,
fastvideo_args, encoder_precision)
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", [])
@@ -264,26 +251,12 @@ 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,
use_cpu_offload))
self._get_all_weights(model_config, model, model_path))
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:
@@ -317,10 +290,10 @@ class ImageEncoderLoader(TextEncoderLoader):
encoder_config = fastvideo_args.pipeline_config.image_encoder_config
encoder_config.update_model_arch(model_config)
target_device = get_local_torch_device()
target_device = get_torch_device()
# TODO(will): add support for other dtypes
return self.load_model(
model_path, encoder_config, target_device, fastvideo_args,
model_path, encoder_config, target_device,
fastvideo_args.pipeline_config.image_encoder_precision)
@@ -373,7 +346,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_local_torch_device())
vae = vae_cls(vae_config).to(get_torch_device())
# Find all safetensors files
safetensors_list = glob.glob(
@@ -432,7 +405,7 @@ class TransformerLoader(ComponentLoader):
"hf_config": hf_config
},
weight_dir_list=safetensors_list,
device=get_local_torch_device(),
device=get_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,
+9 -25
View File
@@ -69,7 +69,6 @@ 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.
@@ -102,12 +101,9 @@ def maybe_load_fsdp_model(
cpu_offload=cpu_offload,
reshard_after_forward=True,
mp_policy=mp_policy,
mesh=device_mesh,
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=pin_cpu_memory)
mesh=device_mesh)
weight_iterator = safetensors_weights_iterator(
weight_dir_list, to_cpu=cpu_offload, async_broadcast=not cpu_offload)
weight_iterator = safetensors_weights_iterator(weight_dir_list)
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
load_model_from_full_model_state_dict(
model,
@@ -130,13 +126,12 @@ def maybe_load_fsdp_model(
def shard_model(
model,
*,
cpu_offload: bool,
reshard_after_forward: bool = True,
mp_policy: Optional[MixedPrecisionPolicy] = MixedPrecisionPolicy(), # noqa
mp_policy: Optional[MixedPrecisionPolicy] = None,
dp_mesh: Optional[DeviceMesh] = None,
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.
@@ -155,28 +150,19 @@ 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.
mesh (Optional[DeviceMesh]): Device mesh to use for FSDP sharding under multiple parallelism.
dp_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(
pin_memory=pin_cpu_memory)
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
# iterating in reverse to start with
# lowest-level modules first
@@ -186,7 +172,7 @@ def shard_model(
for n, m in reversed(list(model.named_modules())):
if any([
shard_condition(n, m)
for shard_condition in fsdp_shard_conditions
for shard_condition in model._fsdp_shard_conditions
]):
fully_shard(m, **fsdp_kwargs)
num_layers_sharded += 1
@@ -195,6 +181,7 @@ 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)
@@ -237,9 +224,6 @@ 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)
+10 -53
View File
@@ -11,11 +11,9 @@ 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__)
@@ -120,77 +118,36 @@ _BAR_FORMAT = "{desc}: {percentage:3.0f}% Completed | {n_fmt}/{total_fmt} [{elap
def safetensors_weights_iterator(
hf_weights_files: List[str],
to_cpu: bool = False,
async_broadcast: bool = False
hf_weights_files: List[str]
) -> Generator[Tuple[str, torch.Tensor], None, None]:
"""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 = []
"""Iterate over the weights in the model safetensor files."""
enable_tqdm = not torch.distributed.is_initialized(
) or torch.distributed.get_rank() == 0
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", device=device) as f:
with safe_open(st_file, framework="pt") as f:
for name in f.keys(): # noqa: SIM118
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))
param = f.get_tensor(name)
yield name, param
if async_broadcast:
for handle in handles:
handle.wait()
def pt_weights_iterator(
hf_weights_files: List[str],
to_cpu: bool = True # default to CPU for text encoder
hf_weights_files: List[str]
) -> Generator[Tuple[str, torch.Tensor], None, None]:
"""Iterate over the weights in the model bin/pt files."""
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
enable_tqdm = not torch.distributed.is_initialized(
) or torch.distributed.get_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=device, weights_only=True)
state = torch.load(bin_file, map_location="cpu", 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_local_torch_device
from fastvideo.v1.distributed import get_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,8 +328,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
# VAE
with torch.autocast("cuda", dtype=torch.float32):
latents = self.get_module("vae").encode(
valid_data["pixel_values"].to(
get_local_torch_device())).mean
valid_data["pixel_values"].to(get_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_local_torch_device
from fastvideo.v1.distributed import get_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_local_torch_device())
self.get_module("vae").to(get_local_torch_device())
self.get_module("image_encoder").to(get_torch_device())
self.get_module("vae").to(get_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_local_torch_device())
dim=0).to(get_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_local_torch_device(), dtype=torch.float32)
video_condition = video_condition.to(device=get_torch_device(),
dtype=torch.float32)
video_conditions.append(video_condition)
video_conditions = torch.cat(video_conditions, dim=0)
+2 -2
View File
@@ -5,7 +5,7 @@ Decoding stage for diffusion pipelines.
import torch
from fastvideo.v1.distributed import get_local_torch_device
from fastvideo.v1.distributed import get_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_local_torch_device())
self.vae = self.vae.to(get_torch_device())
latents = batch.latents
# TODO(will): remove this once we add input/output validation for stages
+3 -4
View File
@@ -12,9 +12,8 @@ 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_local_torch_device,
get_sp_parallel_rank, get_sp_world_size,
get_world_group)
from fastvideo.v1.distributed import (get_sp_parallel_rank, get_sp_world_size,
get_torch_device, get_world_group)
from fastvideo.v1.distributed.communication_op import (
sequence_model_parallel_all_gather)
from fastvideo.v1.fastvideo_args import FastVideoArgs
@@ -193,7 +192,7 @@ class DenoisingStage(PipelineStage):
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=get_local_torch_device(),
device=get_torch_device(),
).to(target_dtype) *
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
is not None else None)
+4 -5
View File
@@ -7,7 +7,7 @@ from typing import Optional
import PIL.Image
import torch
from fastvideo.v1.distributed import get_local_torch_device
from fastvideo.v1.distributed import get_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_local_torch_device())
self.vae = self.vae.to(get_torch_device())
assert batch.height is not None
assert batch.width is not None
@@ -65,8 +65,7 @@ class EncodingStage(PipelineStage):
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=batch.height,
width=batch.width).to(get_local_torch_device(),
dtype=torch.float32)
width=batch.width).to(get_torch_device(), dtype=torch.float32)
image = image.unsqueeze(2)
else:
@@ -79,7 +78,7 @@ class EncodingStage(PipelineStage):
batch.num_frames - 1, batch.height, batch.width)
],
dim=2)
video_condition = video_condition.to(device=get_local_torch_device(),
video_condition = video_condition.to(device=get_torch_device(),
dtype=torch.float32)
# Setup VAE precision
@@ -7,7 +7,7 @@ This module contains implementations of image encoding stages for diffusion pipe
import torch
from fastvideo.v1.distributed import get_local_torch_device
from fastvideo.v1.distributed import get_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_local_torch_device())
self.image_encoder = self.image_encoder.to(get_torch_device())
image = batch.pil_image
image_inputs = self.image_processor(
images=image, return_tensors="pt").to(get_local_torch_device())
images=image, return_tensors="pt").to(get_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_local_torch_device
from fastvideo.v1.distributed import get_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_local_torch_device()
device = get_torch_device()
generator = batch.generator
latents = batch.latents
num_frames = latent_num_frames if latent_num_frames is not None else batch.num_frames
+11 -7
View File
@@ -5,7 +5,9 @@ Prompt encoding stages for diffusion pipelines.
This module contains implementations of prompt encoding stages for diffusion pipelines.
"""
from fastvideo.v1.distributed import get_local_torch_device
import torch
from fastvideo.v1.distributed import get_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
@@ -60,6 +62,8 @@ 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):
@@ -67,9 +71,8 @@ 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_local_torch_device())
text_inputs = tokenizer(
texts, **encoder_config.tokenizer_kwargs).to(get_torch_device())
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
with set_forward_context(current_timestep=0, attn_metadata=None):
@@ -88,8 +91,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_local_torch_device())
negative_text,
**encoder_config.tokenizer_kwargs).to(get_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,
@@ -107,8 +110,9 @@ class TextEncodingStage(PipelineStage):
batch.negative_attention_mask.append(
negative_attention_mask)
if fastvideo_args.text_encoder_offload:
if fastvideo_args.use_cpu_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_local_torch_device
from fastvideo.v1.distributed import get_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_local_torch_device()
device = get_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_local_torch_device
from fastvideo.v1.distributed import get_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_local_torch_device()
target_device = get_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,8 +16,6 @@ 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__)
@@ -68,6 +66,7 @@ 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
@@ -79,20 +78,19 @@ def test_clip_encoder():
logger.info("Model1 has %d parameters", len(params1))
logger.info("Model2 has %d parameters", len(params2))
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()
# 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}")
# 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,8 +15,7 @@ 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"
@@ -63,6 +62,7 @@ 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,28 +77,34 @@ def test_llama_encoder():
# Compare a few key parameters
weight_diffs = []
# check if embed_tokens are the same
device = model1.embed_tokens.weight.device
print(model1.embed_tokens.weight.shape, model2.embed_tokens.weight.shape)
assert torch.allclose(model1.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))
model2.embed_tokens.weight)
weights = [
"layers.{}.input_layernorm.weight",
"layers.{}.post_attention_layernorm.weight"
]
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()
# 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)
tokenizer, _ = load_tokenizer(tokenizer_type="llm",
tokenizer_path=TOKENIZER_PATH,
+16 -12
View File
@@ -4,8 +4,6 @@ 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
@@ -43,13 +41,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,)),
pin_cpu_memory=False)
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH, pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),), text_encoder_precisions=(precision_str,)))
loader = TextEncoderLoader()
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
model2 = model2.to(precision)
# Convert to float16 and move to device
# model2 = model2.to(precision)
model2 = model2.to(device)
model2.eval()
# Sanity check weights between the two models
@@ -66,17 +64,23 @@ 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"]
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight", "shared.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]
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(p1)
assert_close(p1, p2, atol=1e-4, rtol=1e-4)
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)
# Test with some sample prompts
prompts = [
@@ -5,7 +5,7 @@ from pathlib import Path
import pytest
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "1"
NUM_GPUS_PER_NODE = "2"
# 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", "1",
"--tp-size", "1",
"--num-gpus", "1",
"--sp-size", "2",
"--tp-size", "2",
"--num-gpus", "2",
"--height", "768",
"--width", "1280",
"--num-frames", "69",
+2 -2
View File
@@ -69,11 +69,11 @@ def run_ssim_tests():
def run_training_tests():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/v1/tests/training/Vanilla -srP")
@app.function(gpu="H100:1", image=image, timeout=1800, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
@app.function(gpu="H100:2", image=image, timeout=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:1", image=image, timeout=1800)
@app.function(gpu="H100:2", image=image, timeout=1800)
def run_inference_tests_STA():
run_test("pytest ./fastvideo/v1/tests/inference/STA -srP")
@@ -1 +1 @@
{"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}
{"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}
@@ -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 = "1"
NUM_GPUS_PER_NODE = "2"
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", "1",
"--sp_size", "1",
"--tp_size", "1",
"--num_gpus", "2",
"--sp_size", "2",
"--tp_size", "2",
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", "1",
"--hsdp_shard_dim", "2",
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "4",
"--gradient_accumulation_steps", "1",
"--gradient_accumulation_steps", "2",
"--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': 0.5,
'step_time': 1.0,
'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': 1.0,
'avg_step_time': 6.0,
'grad_norm': 0.3,
'step_time': 0.5,
'step_time': 6.0,
'train_loss': 0.0025
}
@@ -80,10 +80,7 @@ def test_hunyuanvideo_distributed():
# Initialize with identical weights
model = initialize_identical_weights(model, seed=42)
shard_model(model, cpu_offload=True,
reshard_after_forward=True,
fsdp_shard_conditions=model._fsdp_shard_conditions
)
shard_model(model, cpu_offload=False, reshard_after_forward=True)
for n, p in chain(model.named_parameters(), model.named_buffers()):
if p.is_meta:
raise RuntimeError(
+10 -13
View File
@@ -24,9 +24,8 @@ 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_local_torch_device, get_sp_group,
get_world_group)
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
get_torch_device, 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
@@ -66,8 +65,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
@@ -177,12 +176,12 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
encoder_attention_mask = batch['text_attention_mask']
infos = batch['info_list']
training_batch.latents = latents.to(get_local_torch_device(),
training_batch.latents = latents.to(get_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_local_torch_device(), dtype=torch.bfloat16)
get_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_local_torch_device(), dtype=torch.bfloat16)
get_torch_device(), dtype=torch.bfloat16)
training_batch.infos = infos
return training_batch
@@ -246,9 +245,8 @@ 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] // patch_size[0],
latents.shape[2] * self.sp_world_size // patch_size[0],
latents.shape[3] // patch_size[1],
latents.shape[4] // patch_size[2]
]
@@ -275,7 +273,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
"encoder_hidden_states":
training_batch.encoder_hidden_states,
"timestep":
training_batch.timesteps.to(get_local_torch_device(),
training_batch.timesteps.to(get_torch_device(),
dtype=torch.bfloat16),
"encoder_attention_mask":
training_batch.encoder_attention_mask,
@@ -552,9 +550,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
prompt_embeds = validation_batch['text_embedding']
prompt_attention_mask = validation_batch['text_attention_mask']
prompt_embeds = prompt_embeds.to(get_local_torch_device())
prompt_attention_mask = prompt_attention_mask.to(
get_local_torch_device())
prompt_embeds = prompt_embeds.to(get_torch_device())
prompt_attention_mask = prompt_attention_mask.to(get_torch_device())
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
@@ -4,11 +4,12 @@ from copy import deepcopy
from typing import Any, Dict
import torch
import torch.distributed
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.dataloader.schema import (
pyarrow_schema_i2v, pyarrow_schema_i2v_validation)
from fastvideo.v1.distributed import get_local_torch_device
from fastvideo.v1.distributed import get_torch_device, get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
@@ -18,6 +19,8 @@ from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
from fastvideo.v1.pipelines.wan.wan_i2v_pipeline import (
WanImageToVideoValidationPipeline)
from fastvideo.v1.training.training_pipeline import TrainingPipeline
from fastvideo.v1.training.training_utils import (shard_latents_across_sp,
clip_grad_norm_while_handling_failing_dtensor_cases)
from fastvideo.v1.utils import is_vsa_available
vsa_available = is_vsa_available()
@@ -85,17 +88,15 @@ class WanI2VTrainingPipeline(TrainingPipeline):
pil_image = batch['pil_image']
infos = batch['info_list']
training_batch.latents = latents.to(get_local_torch_device(),
training_batch.latents = latents.to(get_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_local_torch_device(), dtype=torch.bfloat16)
get_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
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())
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())
training_batch.infos = infos
return training_batch
@@ -114,8 +115,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_local_torch_device(), dtype=torch.bfloat16)
image_latents = training_batch.image_latents.to(get_torch_device(),
dtype=torch.bfloat16)
training_batch.noisy_model_input = torch.cat(
[training_batch.noisy_model_input, image_latents], dim=1)
@@ -134,8 +135,7 @@ 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_local_torch_device(),
dtype=torch.bfloat16)
image_embeds = image_embeds.to(get_torch_device(), dtype=torch.bfloat16)
encoder_hidden_states_image = image_embeds
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
@@ -145,7 +145,7 @@ class WanI2VTrainingPipeline(TrainingPipeline):
"encoder_hidden_states":
training_batch.encoder_hidden_states,
"timestep":
training_batch.timesteps.to(get_local_torch_device(),
training_batch.timesteps.to(get_torch_device(),
dtype=torch.bfloat16),
"encoder_attention_mask":
training_batch.encoder_attention_mask,
@@ -169,9 +169,9 @@ class WanI2VTrainingPipeline(TrainingPipeline):
infos = validation_batch['info_list']
prompt = infos[0]['prompt']
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())
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
clip_features = clip_features.to(get_torch_device())
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
@@ -207,6 +207,36 @@ class WanI2VTrainingPipeline(TrainingPipeline):
)
return batch
def _clip_grad_norm(self, training_batch: TrainingBatch) -> TrainingBatch:
"""Override to add gradient synchronization across SP ranks."""
assert self.training_args is not None
max_grad_norm = self.training_args.max_grad_norm
# CRITICAL FIX: Synchronize gradients across SP ranks before clipping
# Different SP ranks compute different gradients due to different noise patterns
# These gradients must be averaged across SP ranks for stable training
if self.training_args.sp_size > 1:
sp_group = get_sp_group()
for param in self.transformer.parameters():
if param.grad is not None:
# Average gradients across SP ranks
sp_group.all_reduce(param.grad, op=torch.distributed.ReduceOp.AVG)
if max_grad_norm is not None:
model_parts = [self.transformer]
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
foreach=None,
)
assert grad_norm is not float('nan') or grad_norm is not float(
'inf')
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
else:
grad_norm = 0.0
training_batch.grad_norm = grad_norm
return training_batch
def main(args) -> None:
logger.info("Starting training pipeline...")