Compare commits

..
Author SHA1 Message Date
SolitaryThinker d0c53871a4 fix tensor type hint 2025-05-23 14:42:23 -07:00
SolitaryThinker 0d5306f61f update min python to 3.10 2025-05-23 14:42:23 -07:00
115 changed files with 997 additions and 4852 deletions
-2
View File
@@ -77,8 +77,6 @@ jobs:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/transformers/**'
- 'fastvideo/v1/layers/**'
- 'fastvideo/v1/attention/**'
encoder-test:
needs: change-filter
+1 -3
View File
@@ -18,7 +18,6 @@ import os
import re
import sys
from pathlib import Path
from typing import Optional
import requests
@@ -168,8 +167,7 @@ _cached_base: str = ""
_cached_branch: str = ""
def get_repo_base_and_branch(
pr_number: str) -> tuple[Optional[str], Optional[str]]:
def get_repo_base_and_branch(pr_number: str) -> tuple[str | None, str | None]:
global _cached_base, _cached_branch
if _cached_base and _cached_branch:
return _cached_base, _cached_branch
+1 -2
View File
@@ -5,7 +5,6 @@ import itertools
import re
from dataclasses import dataclass, field
from pathlib import Path
from typing import Optional
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
ROOT_DIR_RELATIVE = '../../../..'
@@ -89,7 +88,7 @@ class Example:
generate() -> str: Generates the documentation content.
""" # noqa: E501
path: Path
category: Optional[str] = None
category: str | None = None
main_file: Path = field(init=False)
other_files: list[Path] = field(init=False)
title: str = field(init=False)
-111
View File
@@ -1,111 +0,0 @@
import argparse
import json
import os
import torch
import torch.distributed as dist
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import maybe_download_model, shallow_asdict
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo import PipelineConfig
from fastvideo.v1.pipelines.preprocess_pipeline import PreprocessPipeline
logger = init_logger(__name__)
BASE_MODEL_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
def main(args):
# Assume using torchrun
local_rank = int(os.getenv("RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
init_distributed_environment(world_size=world_size, rank=rank, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
pipeline_config = PipelineConfig.from_pretrained(MODEL_PATH)
kwargs = {
"use_cpu_offload": False,
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
}
pipeline_config_args = shallow_asdict(pipeline_config)
pipeline_config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=MODEL_PATH,
num_gpus=world_size,
device_str="cuda",
**pipeline_config_args,
)
fastvideo_args.check_fastvideo_args()
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
pipeline = PreprocessPipeline(MODEL_PATH, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--preprocess_video_batch_size",
type=int,
default=2,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--preprocess_text_batch_size",
type=int,
default=8,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--max_height", type=int, default=480)
parser.add_argument("--max_width", type=int, default=848)
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--dataset", default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
parser.add_argument("--speed_factor", type=float, default=1.0)
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
# text encoder & vae & diffusion model
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
main(args)
+1 -2
View File
@@ -7,7 +7,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
# from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -38,7 +38,6 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
linear_range=0.5,
):
if linear_quadratic:
raise NotImplementedError("Linear quadratic schedule is not implemented")
linear_steps = int(num_train_timesteps * linear_range)
sigmas = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
@@ -31,7 +31,7 @@ mochi_latents_std = torch.tensor([
mochi_scaling_factor = 1.0
def normalize_dit_input(model_type, latents, args=None):
def normalize_dit_input(model_type, latents):
if model_type == "mochi":
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
@@ -41,16 +41,5 @@ def normalize_dit_input(model_type, latents, args=None):
return latents * 0.476986
elif model_type == "hunyuan":
return latents * 0.476986
elif model_type == "wan":
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
vae_config = WanVAEConfig()
latents_mean = torch.tensor(vae_config.arch_config.latents_mean)
latents_std = 1.0 / torch.tensor(vae_config.arch_config.latents_std)
latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(device=latents.device)
latents_std = latents_std.view(1, -1, 1, 1, 1).to(device=latents.device)
latents = ((latents.float() - latents_mean) * latents_std).to(latents)
return latents
else:
raise NotImplementedError(f"model_type {model_type} not supported")
+6 -8
View File
@@ -3,8 +3,7 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass, fields
from typing import (TYPE_CHECKING, Any, Dict, Generic, Optional, Protocol, Set,
Type, TypeVar)
from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar
if TYPE_CHECKING:
from fastvideo.v1.fastvideo_args import FastVideoArgs
@@ -27,12 +26,12 @@ class AttentionBackend(ABC):
@staticmethod
@abstractmethod
def get_impl_cls() -> Type["AttentionImpl"]:
def get_impl_cls() -> type["AttentionImpl"]:
raise NotImplementedError
@staticmethod
@abstractmethod
def get_metadata_cls() -> Type["AttentionMetadata"]:
def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError
# @staticmethod
@@ -46,7 +45,7 @@ class AttentionBackend(ABC):
@staticmethod
@abstractmethod
def get_builder_cls() -> Type["AttentionMetadataBuilder"]:
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
raise NotImplementedError
@@ -57,8 +56,7 @@ class AttentionMetadata:
current_timestep: int
def asdict_zerocopy(self,
skip_fields: Optional[Set[str]] = None
) -> Dict[str, Any]:
skip_fields: set[str] | None = None) -> dict[str, Any]:
"""Similar to dataclasses.asdict, but avoids deepcopying."""
if skip_fields is None:
skip_fields = set()
@@ -124,7 +122,7 @@ class AttentionImpl(ABC, Generic[T]):
head_size: int,
softmax_scale: float,
causal: bool = False,
num_kv_heads: Optional[int] = None,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
@@ -1,7 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
from typing import List, Optional, Type
import torch
from flash_attn import flash_attn_func as flash_attn_2_func
@@ -28,7 +26,7 @@ class FlashAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> List[int]:
def get_supported_head_sizes() -> list[int]:
return [32, 64, 96, 128, 160, 192, 224, 256]
@staticmethod
@@ -36,15 +34,15 @@ class FlashAttentionBackend(AttentionBackend):
return "FLASH_ATTN"
@staticmethod
def get_impl_cls() -> Type["FlashAttentionImpl"]:
def get_impl_cls() -> type["FlashAttentionImpl"]:
return FlashAttentionImpl
@staticmethod
def get_metadata_cls() -> Type["AttentionMetadata"]:
def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError
@staticmethod
def get_builder_cls() -> Type["AttentionMetadataBuilder"]:
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
raise NotImplementedError
@@ -56,7 +54,7 @@ class FlashAttentionImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
+3 -5
View File
@@ -1,5 +1,3 @@
from typing import List, Optional, Type
import torch
from sageattention import sageattn
@@ -17,7 +15,7 @@ class SageAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> List[int]:
def get_supported_head_sizes() -> list[int]:
return [32, 64, 96, 128, 160, 192, 224, 256]
@staticmethod
@@ -25,7 +23,7 @@ class SageAttentionBackend(AttentionBackend):
return "SAGE_ATTN"
@staticmethod
def get_impl_cls() -> Type["SageAttentionImpl"]:
def get_impl_cls() -> type["SageAttentionImpl"]:
return SageAttentionImpl
# @staticmethod
@@ -41,7 +39,7 @@ class SageAttentionImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
+3 -5
View File
@@ -1,5 +1,3 @@
from typing import List, Optional, Type
import torch
from fastvideo.v1.attention.backends.abstract import (
@@ -16,7 +14,7 @@ class SDPABackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> List[int]:
def get_supported_head_sizes() -> list[int]:
return [32, 64, 96, 128, 160, 192, 224, 256]
@staticmethod
@@ -24,7 +22,7 @@ class SDPABackend(AttentionBackend):
return "SDPA"
@staticmethod
def get_impl_cls() -> Type["SDPAImpl"]:
def get_impl_cls() -> type["SDPAImpl"]:
return SDPAImpl
# @staticmethod
@@ -40,7 +38,7 @@ class SDPAImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
@@ -1,6 +1,5 @@
import json
from dataclasses import dataclass
from typing import List, Optional, Type
import torch
from einops import rearrange
@@ -20,7 +19,7 @@ logger = init_logger(__name__)
# TODO(will-refactor): move this to a utils file
def dict_to_3d_list(mask_strategy) -> List[List[List[Optional[torch.Tensor]]]]:
def dict_to_3d_list(mask_strategy) -> list[list[list[torch.Tensor | None]]]:
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
max_timesteps_idx = max(
@@ -58,7 +57,7 @@ class SlidingTileAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> List[int]:
def get_supported_head_sizes() -> list[int]:
# TODO(will-refactor): check this
return [32, 64, 96, 128, 160, 192, 224, 256]
@@ -67,15 +66,15 @@ class SlidingTileAttentionBackend(AttentionBackend):
return "SLIDING_TILE_ATTN"
@staticmethod
def get_impl_cls() -> Type["SlidingTileAttentionImpl"]:
def get_impl_cls() -> type["SlidingTileAttentionImpl"]:
return SlidingTileAttentionImpl
@staticmethod
def get_metadata_cls() -> Type["SlidingTileAttentionMetadata"]:
def get_metadata_cls() -> type["SlidingTileAttentionMetadata"]:
return SlidingTileAttentionMetadata
@staticmethod
def get_builder_cls() -> Type["SlidingTileAttentionMetadataBuilder"]:
def get_builder_cls() -> type["SlidingTileAttentionMetadataBuilder"]:
return SlidingTileAttentionMetadataBuilder
@@ -110,7 +109,7 @@ class SlidingTileAttentionImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
+12 -14
View File
@@ -1,7 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional, Tuple
import torch
import torch.nn as nn
@@ -22,11 +20,11 @@ class DistributedAttention(nn.Module):
def __init__(self,
num_heads: int,
head_size: int,
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
num_kv_heads: int | None = None,
softmax_scale: float | None = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
supported_attention_backends: tuple[_Backend, ...]
| None = None,
prefix: str = "",
**extra_impl_args) -> None:
super().__init__()
@@ -62,10 +60,10 @@ class DistributedAttention(nn.Module):
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
replicated_q: Optional[torch.Tensor] = None,
replicated_k: Optional[torch.Tensor] = None,
replicated_v: Optional[torch.Tensor] = None,
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
replicated_q: torch.Tensor | None = None,
replicated_k: torch.Tensor | None = None,
replicated_v: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Forward pass for distributed attention.
Args:
@@ -141,11 +139,11 @@ class LocalAttention(nn.Module):
def __init__(self,
num_heads: int,
head_size: int,
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
num_kv_heads: int | None = None,
softmax_scale: float | None = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
supported_attention_backends: tuple[_Backend, ...]
| None = None,
**extra_impl_args) -> None:
super().__init__()
if softmax_scale is None:
+14 -13
View File
@@ -2,9 +2,10 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/selector.py
import os
from collections.abc import Generator
from contextlib import contextmanager
from functools import cache
from typing import Generator, Optional, Tuple, Type, cast
from typing import cast
import torch
@@ -17,7 +18,7 @@ from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
logger = init_logger(__name__)
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
def backend_name_to_enum(backend_name: str) -> _Backend | None:
"""
Convert a string backend name to a _Backend enum value.
@@ -31,7 +32,7 @@ def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
None
def get_env_variable_attn_backend() -> Optional[_Backend]:
def get_env_variable_attn_backend() -> _Backend | None:
'''
Get the backend override specified by the FastVideo attention
backend environment variable, if one is specified.
@@ -53,10 +54,10 @@ def get_env_variable_attn_backend() -> Optional[_Backend]:
#
# THIS SELECTION TAKES PRECEDENCE OVER THE
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
forced_attn_backend: Optional[_Backend] = None
forced_attn_backend: _Backend | None = None
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
def global_force_attn_backend(attn_backend: _Backend | None) -> None:
'''
Force all attention operations to use a specified backend.
@@ -71,7 +72,7 @@ def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
forced_attn_backend = attn_backend
def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_global_forced_attn_backend() -> _Backend | None:
'''
Get the currently-forced choice of attention backend,
or None if auto-selection is currently enabled.
@@ -82,8 +83,8 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
) -> Type[AttentionBackend]:
supported_attention_backends: tuple[_Backend, ...] | None = None,
) -> type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype,
supported_attention_backends)
@@ -92,8 +93,8 @@ def get_attn_backend(
def _cached_get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
) -> Type[AttentionBackend]:
supported_attention_backends: tuple[_Backend, ...] | None = None,
) -> type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
#
@@ -102,13 +103,13 @@ def _cached_get_attn_backend(
if not supported_attention_backends:
raise ValueError("supported_attention_backends is empty")
selected_backend = None
backend_by_global_setting: Optional[_Backend] = (
backend_by_global_setting: _Backend | None = (
get_global_forced_attn_backend())
if backend_by_global_setting is not None:
selected_backend = backend_by_global_setting
else:
# Check the environment variable and override if specified
backend_by_env_var: Optional[str] = envs.FASTVIDEO_ATTENTION_BACKEND
backend_by_env_var: str | None = envs.FASTVIDEO_ATTENTION_BACKEND
if backend_by_env_var is not None:
selected_backend = backend_name_to_enum(backend_by_env_var)
@@ -120,7 +121,7 @@ def _cached_get_attn_backend(
if not attention_cls:
raise ValueError(
f"Invalid attention backend for {current_platform.device_name}")
return cast(Type[AttentionBackend], resolve_obj_by_qualname(attention_cls))
return cast(type[AttentionBackend], resolve_obj_by_qualname(attention_cls))
@contextmanager
+3 -3
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field, fields
from typing import Any, Dict
from typing import Any
from fastvideo.v1.logger import init_logger
@@ -41,7 +41,7 @@ class ModelConfig:
self.__dict__.update(state)
# This should be used only when loading from transformers/diffusers
def update_model_arch(self, source_model_dict: Dict[str, Any]) -> None:
def update_model_arch(self, source_model_dict: dict[str, Any]) -> None:
arch_config = self.arch_config
valid_fields = {f.name for f in fields(arch_config)}
@@ -55,7 +55,7 @@ class ModelConfig:
if hasattr(arch_config, "__post_init__"):
arch_config.__post_init__()
def update_model_config(self, source_model_dict: Dict[str, Any]) -> None:
def update_model_config(self, source_model_dict: dict[str, Any]) -> None:
assert "arch_config" not in source_model_dict, "Source model config shouldn't contain arch_config."
valid_fields = {f.name for f in fields(self)}
+3 -3
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Any, Optional, Tuple
from typing import Any
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
@@ -11,7 +11,7 @@ class DiTArchConfig(ArchConfig):
_fsdp_shard_conditions: list = field(default_factory=list)
_compile_conditions: list = field(default_factory=list)
_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: Tuple[_Backend,
_supported_attention_backends: tuple[_Backend,
...] = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_ATTN,
_Backend.FLASH_ATTN,
@@ -32,7 +32,7 @@ class DiTConfig(ModelConfig):
# FastVideoDiT-specific parameters
prefix: str = ""
quant_config: Optional[QuantizationConfig] = None
quant_config: QuantizationConfig | None = None
@staticmethod
def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any:
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
import torch
@@ -156,9 +155,9 @@ class HunyuanVideoArchConfig(DiTArchConfig):
num_layers: int = 20
num_single_layers: int = 40
num_refiner_layers: int = 2
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56)
rope_axes_dim: tuple[int, int, int] = (16, 56, 56)
guidance_embeds: bool = False
dtype: Optional[torch.dtype] = None
dtype: torch.dtype | None = None
text_embed_dim: int = 4096
pooled_projection_dim: int = 768
rope_theta: int = 256
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import List, Optional, Tuple, Union
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -40,17 +39,17 @@ class StepVideoArchConfig(DiTArchConfig):
num_attention_heads: int = 48
attention_head_dim: int = 128
in_channels: int = 64
out_channels: Optional[int] = 64
out_channels: int | None = 64
num_layers: int = 48
dropout: float = 0.0
patch_size: int = 1
norm_type: str = "ada_norm_single"
norm_elementwise_affine: bool = False
norm_eps: float = 1e-6
caption_channels: Optional[Union[int, List[int], Tuple[int, ...]]] = field(
caption_channels: int | list[int] | tuple[int, ...] | None = field(
default_factory=lambda: [6144, 1024])
attention_type: Optional[str] = "torch"
use_additional_conditions: Optional[bool] = False
attention_type: str | None = "torch"
use_additional_conditions: bool | None = False
def __post_init__(self):
self.hidden_size = self.num_attention_heads * self.attention_head_dim
+3 -4
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -52,7 +51,7 @@ class WanVideoArchConfig(DiTArchConfig):
r"blocks.\1.self_attn_residual_norm.norm.\2",
})
patch_size: Tuple[int, int, int] = (1, 2, 2)
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len = 512
num_attention_heads: int = 40
attention_head_dim: int = 128
@@ -65,8 +64,8 @@ class WanVideoArchConfig(DiTArchConfig):
cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6
image_dim: Optional[int] = None
added_kv_proj_dim: Optional[int] = None
image_dim: int | None = None
added_kv_proj_dim: int | None = None
rope_max_seq_len: int = 1024
def __post_init__(self):
+11 -11
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Tuple
from typing import Any
import torch
@@ -10,8 +10,8 @@ from fastvideo.v1.platforms import _Backend
@dataclass
class EncoderArchConfig(ArchConfig):
architectures: List[str] = field(default_factory=lambda: [])
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
architectures: list[str] = field(default_factory=lambda: [])
_supported_attention_backends: tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
output_hidden_states: bool = False
use_return_dict: bool = True
@@ -32,7 +32,7 @@ class TextEncoderArchConfig(EncoderArchConfig):
scalable_attention: bool = True
tie_word_embeddings: bool = False
tokenizer_kwargs: Dict[str, Any] = field(default_factory=dict)
tokenizer_kwargs: dict[str, Any] = field(default_factory=dict)
def __post_init__(self) -> None:
self.tokenizer_kwargs = {
@@ -49,11 +49,11 @@ class ImageEncoderArchConfig(EncoderArchConfig):
@dataclass
class BaseEncoderOutput:
last_hidden_state: Optional[torch.FloatTensor] = None
pooler_output: Optional[torch.FloatTensor] = None
hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
attentions: Optional[Tuple[torch.FloatTensor, ...]] = None
attention_mask: Optional[torch.Tensor] = None
last_hidden_state: torch.FloatTensor | None = None
pooler_output: torch.FloatTensor | None = None
hidden_states: tuple[torch.FloatTensor, ...] | None = None
attentions: tuple[torch.FloatTensor, ...] | None = None
attention_mask: torch.Tensor | None = None
@dataclass
@@ -61,8 +61,8 @@ class EncoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=EncoderArchConfig)
prefix: str = ""
quant_config: Optional[QuantizationConfig] = None
lora_config: Optional[Any] = None
quant_config: QuantizationConfig | None = None
lora_config: Any | None = None
@dataclass
+4 -5
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
ImageEncoderConfig,
@@ -51,8 +50,8 @@ class CLIPTextConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(
default_factory=CLIPTextArchConfig)
num_hidden_layers_override: Optional[int] = None
require_post_norm: Optional[bool] = None
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
prefix: str = "clip"
@@ -61,6 +60,6 @@ class CLIPVisionConfig(ImageEncoderConfig):
arch_config: ImageEncoderArchConfig = field(
default_factory=CLIPVisionArchConfig)
num_hidden_layers_override: Optional[int] = None
require_post_norm: Optional[bool] = None
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
prefix: str = "clip"
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
@@ -12,7 +11,7 @@ class LlamaArchConfig(TextEncoderArchConfig):
intermediate_size: int = 11008
num_hidden_layers: int = 32
num_attention_heads: int = 32
num_key_value_heads: Optional[int] = None
num_key_value_heads: int | None = None
hidden_act: str = "silu"
max_position_embeddings: int = 2048
initializer_range: float = 0.02
@@ -24,11 +23,11 @@ class LlamaArchConfig(TextEncoderArchConfig):
pretraining_tp: int = 1
tie_word_embeddings: bool = False
rope_theta: float = 10000.0
rope_scaling: Optional[float] = None
rope_scaling: float | None = None
attention_bias: bool = False
attention_dropout: float = 0.0
mlp_bias: bool = False
head_dim: Optional[int] = None
head_dim: int | None = None
hidden_state_skip_layer: int = 2
text_len: int = 256
+1 -2
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
@@ -12,7 +11,7 @@ class T5ArchConfig(TextEncoderArchConfig):
d_kv: int = 64
d_ff: int = 2048
num_layers: int = 6
num_decoder_layers: Optional[int] = None
num_decoder_layers: int | None = None
num_heads: int = 8
relative_attention_num_buckets: int = 32
relative_attention_max_distance: int = 128
+2 -2
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Any, Union
from typing import Any
import torch
@@ -9,7 +9,7 @@ from fastvideo.v1.utils import StoreBoolean
@dataclass
class VAEArchConfig(ArchConfig):
scaling_factor: Union[float, torch.tensor] = 0
scaling_factor: float | torch.Tensor = 0
temporal_compression_ratio: int = 4
spatial_compression_ratio: int = 8
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Tuple
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
@@ -9,19 +8,19 @@ class HunyuanVAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 16
down_block_types: Tuple[str, ...] = (
down_block_types: tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
)
up_block_types: Tuple[str, ...] = (
up_block_types: tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
)
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512)
block_out_channels: tuple[int, ...] = (128, 256, 512, 512)
layers_per_block: int = 2
act_fn: str = "silu"
norm_num_groups: int = 32
+7 -8
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Tuple
import torch
@@ -10,12 +9,12 @@ from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
class WanVAEArchConfig(VAEArchConfig):
base_dim: int = 96
z_dim: int = 16
dim_mult: Tuple[int, ...] = (1, 2, 4, 4)
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
attn_scales: Tuple[float, ...] = ()
temperal_downsample: Tuple[bool, ...] = (False, True, True)
attn_scales: tuple[float, ...] = ()
temperal_downsample: tuple[bool, ...] = (False, True, True)
dropout: float = 0.0
latents_mean: Tuple[float, ...] = (
latents_mean: tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
@@ -33,7 +32,7 @@ class WanVAEArchConfig(VAEArchConfig):
0.2503,
-0.2921,
)
latents_std: Tuple[float, ...] = (
latents_std: tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
@@ -55,9 +54,9 @@ class WanVAEArchConfig(VAEArchConfig):
spatial_compression_ratio = 8
def __post_init__(self):
self.scaling_factor: torch.tensor = 1.0 / torch.tensor(
self.scaling_factor: torch.Tensor = 1.0 / torch.tensor(
self.latents_std).view(1, self.z_dim, 1, 1, 1)
self.shift_factor: torch.tensor = torch.tensor(self.latents_mean).view(
self.shift_factor: torch.Tensor = torch.tensor(self.latents_mean).view(
1, self.z_dim, 1, 1, 1)
+13 -11
View File
@@ -1,6 +1,7 @@
import json
from collections.abc import Callable
from dataclasses import asdict, dataclass, field, fields
from typing import Any, Callable, Dict, Optional, Tuple, cast
from typing import Any, cast
import torch
@@ -17,7 +18,7 @@ def preprocess_text(prompt: str) -> str:
return prompt
def postprocess_text(output: BaseEncoderOutput) -> torch.tensor:
def postprocess_text(output: BaseEncoderOutput) -> torch.Tensor:
raise NotImplementedError
@@ -26,7 +27,7 @@ class PipelineConfig:
"""Base configuration for all pipeline architectures."""
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
flow_shift: float | None = None
use_cpu_offload: bool = False
disable_autocast: bool = False
@@ -43,18 +44,18 @@ class PipelineConfig:
dit_config: DiTConfig = field(default_factory=DiTConfig)
# Text encoder configuration
text_encoder_precisions: Tuple[str, ...] = field(
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp16", ))
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), ))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(postprocess_text, ))
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
mask_strategy_file_path: str | None = None
# Compilation
enable_torch_compile: bool = False
@@ -107,7 +108,7 @@ class PipelineConfig:
input_pipeline_dict = json.load(f)
self.update_pipeline_config(input_pipeline_dict)
def update_pipeline_config(self, source_pipeline_dict: Dict[str,
def update_pipeline_config(self, source_pipeline_dict: dict[str,
Any]) -> None:
for f in fields(self):
key = f.name
@@ -123,8 +124,9 @@ class PipelineConfig:
assert len(current_value) == len(
new_value
), "Users shouldn't delete or add text encoder config objects in your json"
for target_config, source_config in zip(
current_value, new_value):
for target_config, source_config in zip(current_value,
new_value,
strict=False):
target_config.update_model_config(source_config)
else:
setattr(self, key, new_value)
+12 -11
View File
@@ -1,5 +1,6 @@
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Callable, Tuple, TypedDict
from typing import TypedDict
import torch
@@ -35,11 +36,11 @@ def llama_preprocess_text(prompt: str) -> str:
return prompt_template_video["template"].format(prompt)
def llama_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
def llama_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
hidden_state_skip_layer = 2
assert outputs.hidden_states is not None
hidden_states: Tuple[torch.Tensor, ...] = outputs.hidden_states
last_hidden_state: torch.tensor = hidden_states[-(hidden_state_skip_layer +
hidden_states: tuple[torch.Tensor, ...] = outputs.hidden_states
last_hidden_state: torch.Tensor = hidden_states[-(hidden_state_skip_layer +
1)]
crop_start = prompt_template_video.get("crop_start", -1)
last_hidden_state = last_hidden_state[:, crop_start:]
@@ -50,8 +51,8 @@ def clip_preprocess_text(prompt: str) -> str:
return prompt
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
pooler_output: torch.tensor = outputs.pooler_output
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
pooler_output: torch.Tensor = outputs.pooler_output
return pooler_output
@@ -72,19 +73,19 @@ class HunyuanConfig(PipelineConfig):
use_cpu_offload: bool = True
# Text encoding stage
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LlamaConfig(), CLIPTextConfig()))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (llama_preprocess_text, clip_preprocess_text))
postprocess_text_funcs: Tuple[
Callable[[BaseEncoderOutput], torch.tensor],
postprocess_text_funcs: tuple[
Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(llama_postprocess_text, clip_postprocess_text))
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field(
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp16"))
def __post_init__(self):
+5 -5
View File
@@ -1,7 +1,7 @@
"""Registry for pipeline weight-specific configurations."""
import os
from typing import Callable, Dict, Optional, Type
from collections.abc import Callable
from fastvideo.v1.configs.pipelines.base import PipelineConfig
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
@@ -18,7 +18,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
WEIGHT_CONFIG_REGISTRY: dict[str, type[PipelineConfig]] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
@@ -30,7 +30,7 @@ WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
}
# For determining pipeline type from model ID
PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"hunyuan": lambda id: "hunyuan" in id.lower(),
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
@@ -39,7 +39,7 @@ PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
}
# Fallback configs when exact match isn't found but architecture is detected
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
"hunyuan":
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
"wanpipeline":
@@ -51,7 +51,7 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
def get_pipeline_config_cls_for_name(
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
pipeline_name_or_path: str) -> type[PipelineConfig] | None:
"""Get the appropriate config class for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
+11 -9
View File
@@ -1,5 +1,5 @@
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Callable, Tuple
import torch
@@ -11,13 +11,15 @@ from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
mask: torch.tensor = outputs.attention_mask
hidden_state: torch.tensor = outputs.last_hidden_state
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
mask: torch.Tensor = outputs.attention_mask
hidden_state: torch.Tensor = outputs.last_hidden_state
seq_lens = mask.gt(0).sum(dim=1).long()
assert torch.isnan(hidden_state).sum() == 0
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens)]
prompt_embeds_tensor: torch.tensor = torch.stack([
prompt_embeds = [
u[:v] for u, v in zip(hidden_state, seq_lens, strict=False)
]
prompt_embeds_tensor: torch.Tensor = torch.stack([
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
for u in prompt_embeds
],
@@ -44,16 +46,16 @@ class WanT2V480PConfig(PipelineConfig):
flow_shift: int = 3
# Text encoding stage
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (T5Config(), ))
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(t5_postprocess_text, ))
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field(
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp32", ))
# WanConfig-specific added parameters
+6 -6
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Union
from typing import Any
from fastvideo.v1.logger import init_logger
@@ -15,12 +15,12 @@ class SamplingParam:
data_type: str = "video"
# Image inputs
image_path: Optional[str] = None
image_path: str | None = None
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
negative_prompt: Optional[str] = None
prompt_path: Optional[str] = None
prompt: str | list[str] | None = None
negative_prompt: str | None = None
prompt_path: str | None = None
output_path: str = "outputs/"
# Batch info
@@ -53,7 +53,7 @@ class SamplingParam:
if self.prompt_path and not self.prompt_path.endswith(".txt"):
raise ValueError("prompt_path must be a txt file")
def update(self, source_dict: Dict[str, Any]) -> None:
def update(self, source_dict: dict[str, Any]) -> None:
for key, value in source_dict.items():
if hasattr(self, key):
setattr(self, key, value)
+6 -6
View File
@@ -1,5 +1,6 @@
import os
from typing import Any, Callable, Dict, Optional
from collections.abc import Callable
from typing import Any
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
@@ -14,7 +15,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
SAMPLING_PARAM_REGISTRY: Dict[str, Any] = {
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
@@ -26,7 +27,7 @@ SAMPLING_PARAM_REGISTRY: Dict[str, Any] = {
}
# For determining pipeline type from model ID
SAMPLING_PARAM_DETECTOR: Dict[str, Callable[[str], bool]] = {
SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
"hunyuan": lambda id: "hunyuan" in id.lower(),
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
@@ -35,7 +36,7 @@ SAMPLING_PARAM_DETECTOR: Dict[str, Callable[[str], bool]] = {
}
# Fallback configs when exact match isn't found but architecture is detected
SAMPLING_FALLBACK_PARAM: Dict[str, Any] = {
SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
"hunyuan":
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
"wanpipeline":
@@ -46,8 +47,7 @@ SAMPLING_FALLBACK_PARAM: Dict[str, Any] = {
}
def get_sampling_param_cls_for_name(
pipeline_name_or_path: str) -> Optional[Any]:
def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
"""Get the appropriate sampling param for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
-39
View File
@@ -1,39 +0,0 @@
from torchvision import transforms
from torchvision.transforms import Lambda
from transformers import AutoTokenizer
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
def getdataset(args, start_idx=0) -> T2V_dataset:
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
resize_topcrop = [
CenterCropResizeVideo((args.max_height, args.max_width), top_crop=True),
]
resize = [
CenterCropResizeVideo((args.max_height, args.max_width)),
]
transform = transforms.Compose([
# Normalize255(),
*resize,
])
transform_topcrop = transforms.Compose([
Normalize255(),
*resize_topcrop,
norm_fun,
])
# tokenizer = AutoTokenizer.from_pretrained("/storage/ongoing/new/Open-Sora-Plan/cache_dir/mt5-xxl", cache_dir=args.cache_dir)
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name,
cache_dir=args.cache_dir)
if args.dataset == "t2v":
return T2V_dataset(args,
transform=transform,
temporal_sample=temporal_sample,
tokenizer=tokenizer,
transform_topcrop=transform_topcrop,
start_idx=start_idx)
raise NotImplementedError(args.dataset)
-44
View File
@@ -1,44 +0,0 @@
# schema.py
"""
Unified data schema and format for saving and loading image/video data after
preprocessing.
It uses apache arrow in-memory format that can be consumed by modern data
frameworks that can handle parquet or lance file.
"""
import pyarrow as pa
pyarrow_schema = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
-129
View File
@@ -1,129 +0,0 @@
import json
import os
import random
import torch
from torch.utils.data import Dataset
class LatentDataset(Dataset):
def __init__(
self,
json_path,
num_latent_t,
cfg_rate,
) -> None:
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
self.json_path = json_path
self.cfg_rate = cfg_rate
self.datase_dir_path = os.path.dirname(json_path)
self.video_dir = os.path.join(self.datase_dir_path, "video")
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
self.prompt_embed_dir = os.path.join(self.datase_dir_path,
"prompt_embed")
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path,
"prompt_attention_mask")
with open(self.json_path) as f:
self.data_anno = json.load(f)
# json.load(f) already keeps the order
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
self.num_latent_t = num_latent_t
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
self.uncond_prompt_mask = torch.zeros(256).bool()
self.lengths = [
data_item.get("length", 1) for data_item in self.data_anno
]
def __getitem__(self, idx):
latent_file = self.data_anno[idx]["latent_path"]
prompt_embed_file = self.data_anno[idx]["prompt_embed_path"]
prompt_attention_mask_file = self.data_anno[idx][
"prompt_attention_mask"]
# load
latent = torch.load(
os.path.join(self.latent_dir, latent_file),
map_location="cpu",
weights_only=True,
)
latent = latent.squeeze(0)[:, -self.num_latent_t:]
if random.random() < self.cfg_rate:
prompt_embed = self.uncond_prompt_embed
prompt_attention_mask = self.uncond_prompt_mask
else:
prompt_embed = torch.load(
os.path.join(self.prompt_embed_dir, prompt_embed_file),
map_location="cpu",
weights_only=True,
)
prompt_attention_mask = torch.load(
os.path.join(self.prompt_attention_mask_dir,
prompt_attention_mask_file),
map_location="cpu",
weights_only=True,
)
return latent, prompt_embed, prompt_attention_mask
def __len__(self):
return len(self.data_anno)
def latent_collate_function(batch):
# return latent, prompt, latent_attn_mask, text_attn_mask
# latent_attn_mask: # b t h w
# text_attn_mask: b 1 l
# needs to check if the latent/prompt' size and apply padding & attn mask
latents, prompt_embeds, prompt_attention_masks = zip(*batch)
# calculate max shape
max_t = max([latent.shape[1] for latent in latents])
max_h = max([latent.shape[2] for latent in latents])
max_w = max([latent.shape[3] for latent in latents])
# padding
latent_list: list[torch.Tensor] = [
torch.nn.functional.pad(
latent,
(
0,
max_t - latent.shape[1],
0,
max_h - latent.shape[2],
0,
max_w - latent.shape[3],
),
) for latent in latents
]
# attn mask
latent_attn_mask = torch.ones(len(latent_list), max_t, max_h, max_w)
# set to 0 if padding
for i, latent in enumerate(latent_list):
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
prompt_embeds = torch.stack(prompt_embeds, dim=0)
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
latents = torch.stack(latent_list, dim=0)
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
if __name__ == "__main__":
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt",
num_latent_t=28,
cfg_rate=0.0)
dataloader = torch.utils.data.DataLoader(dataset,
batch_size=2,
shuffle=False,
collate_fn=latent_collate_function)
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
print(
latent.shape,
prompt_embed.shape,
latent_attn_mask.shape,
prompt_attention_mask.shape,
)
import pdb
pdb.set_trace()
-369
View File
@@ -1,369 +0,0 @@
import argparse
import json
import os
import random
import time
from collections import defaultdict
from typing import Any, Dict, List
import numpy as np
import pyarrow.parquet as pq
import torch
import tqdm
from einops import rearrange
from torch import distributed as dist
from torch.utils.data import Dataset
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
get_sp_group)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class ParquetVideoTextDataset(Dataset):
"""Efficient loader for video-text data from a directory of Parquet files."""
def __init__(self,
path: str,
batch_size: int = 1024,
rank: int = 0,
world_size: int = 1,
cfg_rate: float = 0.0,
num_latent_t: int = 2,
seed: int = 0):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.rank = rank
self.local_rank = get_sequence_model_parallel_rank()
self.sp_world_size = world_size
self.world_size = int(os.getenv("WORLD_SIZE", 1))
self.cfg_rate = cfg_rate
self.num_latent_t = num_latent_t
self.local_indices = None
self.plan_output_dir = os.path.join(self.path, "data_plan.json")
ranks = get_sp_group().ranks
group_ranks: List[List] = [[] for _ in range(self.world_size)]
torch.distributed.all_gather_object(group_ranks, ranks)
if rank == 0:
# If a plan already exists, then skip creating a new plan
# This will be useful when resume training
if os.path.exists(self.plan_output_dir):
print(f"Using existing plan from {self.plan_output_dir}")
return
# Find all parquet files recursively, and record num_rows for each file
print(f"Scanning for parquet files in {self.path}")
metadatas = []
for root, _, files in os.walk(self.path):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.join(root, file)
num_rows = pq.ParquetFile(file_path).metadata.num_rows
for row_idx in range(num_rows):
metadatas.append((file_path, row_idx))
# Generate the plan that distribute rows among workers
random.seed(seed)
random.shuffle(metadatas)
# Get all sp groups
# e.g. if num_gpus = 4, sp_size = 2
# group_ranks = [(0, 1), (2, 3)]
# We will assign the same batches of data to ranks in the same sp group, and we'll assign different batches to ranks in different sp groups
# e.g. plan = {0: [row 1, row 4], 1: [row 1, row 4], 2: [row 2, row 3], 3: [row 2, row 3]}
group_ranks_list: List[Any] = list(
set(tuple(r) for r in group_ranks))
num_sp_groups = len(group_ranks_list)
plan = defaultdict(list)
for idx, metadata in enumerate(metadatas):
sp_group_idx = idx % num_sp_groups
for global_rank in group_ranks_list[sp_group_idx]:
plan[global_rank].append(metadata)
with open(self.plan_output_dir, "w") as f:
json.dump(plan, f)
def __len__(self):
if self.local_indices is None:
try:
with open(self.plan_output_dir) as f:
plan = json.load(f)
self.local_indices = plan[str(self.rank)]
except Exception as err:
raise Exception(
"The data plan hasn't been created yet") from err
assert self.local_indices is not None
return len(self.local_indices)
def __getitem__(self, idx):
if self.local_indices is None:
try:
with open(self.plan_output_dir) as f:
plan = json.load(f)
self.local_indices = plan[self.rank]
except Exception as err:
raise Exception(
"The data plan hasn't been created yet") from err
assert self.local_indices is not None
file_path, row_idx = self.local_indices[idx]
parquet_file = pq.ParquetFile(file_path)
# Calculate the row group to read into memory and the local idx
# This way we can avoid reading in the entire parquet file
cumulative = 0
for i in range(parquet_file.num_row_groups):
num_rows = parquet_file.metadata.row_group(i).num_rows
if cumulative + num_rows > idx:
row_group_index = i
local_index = idx - cumulative
break
cumulative += num_rows
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
row_dict = {k: v[local_index] for k, v in row_group.items()}
del row_group
processed = self._process_row(row_dict)
lat, emb, mask, info = processed["latents"], processed[
"embeddings"], processed["masks"], processed["info"]
if lat.numel() == 0: # Validation parquet
return lat, emb, mask, info
else:
lat = lat[:, -self.num_latent_t:]
if self.sp_world_size > 1:
lat = rearrange(lat,
"t (n s) h w -> t n s h w",
n=self.sp_world_size).contiguous()
lat = lat[:, self.local_rank, :, :, :]
return lat, emb, mask, info
def _process_row(self, row) -> Dict[str, Any]:
"""Process a PyArrow batch into tensors."""
vae_latent_bytes = row["vae_latent_bytes"]
vae_latent_shape = row["vae_latent_shape"]
text_embedding_bytes = row["text_embedding_bytes"]
text_embedding_shape = row["text_embedding_shape"]
text_attention_mask_bytes = row["text_attention_mask_bytes"]
text_attention_mask_shape = row["text_attention_mask_shape"]
# Process latent
if not vae_latent_shape: # No VAE latent is stored. Split is validation
lat = np.array([])
else:
lat = np.frombuffer(vae_latent_bytes,
dtype=np.float32).reshape(vae_latent_shape)
# Make array writable
lat = np.copy(lat)
if random.random() < self.cfg_rate:
emb = np.zeros((512, 4096), dtype=np.float32)
else:
emb = np.frombuffer(text_embedding_bytes,
dtype=np.float32).reshape(text_embedding_shape)
# Make array writable
emb = np.copy(emb)
if emb.shape[0] < 512:
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
padded_emb[:emb.shape[0], :] = emb
emb = padded_emb
elif emb.shape[0] > 512:
emb = emb[:512, :]
# Process mask
if len(text_attention_mask_bytes) > 0 and len(
text_attention_mask_shape) > 0:
msk = np.frombuffer(text_attention_mask_bytes,
dtype=np.uint8).astype(np.bool_)
msk = msk.reshape(1, -1)
# Make array writable
msk = np.copy(msk)
if msk.shape[1] < 512:
padded_msk = np.zeros((1, 512), dtype=np.bool_)
padded_msk[:, :msk.shape[1]] = msk
msk = padded_msk
elif msk.shape[1] > 512:
msk = msk[:, :512]
else:
msk = np.ones((1, 512), dtype=np.bool_)
# Collect metadata
info = {
"width": row["width"],
"height": row["height"],
"num_frames": row["num_frames"],
"duration_sec": row["duration_sec"],
"fps": row["fps"],
"file_name": row["file_name"],
"caption": row["caption"],
}
return {
"latents": torch.from_numpy(lat),
"embeddings": torch.from_numpy(emb),
"masks": torch.from_numpy(msk),
"info": info
}
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description='Benchmark Parquet dataset loading speed')
parser.add_argument('--path',
type=str,
default="your/dataset/path",
help='Path to Parquet dataset')
parser.add_argument('--batch_size',
type=int,
default=4,
help='Batch size for DataLoader')
parser.add_argument('--num_batches',
type=int,
default=100,
help='Number of batches to benchmark')
parser.add_argument('--vae_debug', action="store_true")
args = parser.parse_args()
# Initialize distributed training
local_rank = int(os.environ.get("LOCAL_RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
rank = int(os.environ.get("RANK", 0))
# Initialize CUDA device first
if torch.cuda.is_available():
torch.cuda.set_device(local_rank)
device = torch.device(f"cuda:{local_rank}")
else:
device = torch.device("cpu")
# Initialize distributed training
if world_size > 1:
dist.init_process_group(backend="nccl",
init_method="env://",
world_size=world_size,
rank=rank)
print(
f"Initialized process: rank={rank}, local_rank={local_rank}, world_size={world_size}, device={device}"
)
# Create dataset
dataset = ParquetVideoTextDataset(
args.path,
batch_size=args.batch_size,
rank=rank,
world_size=world_size,
)
# Create DataLoader with proper settings
dataloader = StatefulDataLoader(
dataset,
batch_size=args.batch_size,
num_workers=1, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=True)
# Example of how to load dataloader state
# if os.path.exists("/workspace/FastVideo/dataloader_state.pt"):
# dataloader_state = torch.load("/workspace/FastVideo/dataloader_state.pt")
# dataloader.load_state_dict(dataloader_state[rank])
# Warm-up with synchronization
if rank == 0:
print("Warming up...")
for i, (latents, embeddings, masks, infos) in enumerate(dataloader):
# Example of how to save dataloader state
# if i == 30:
# dist.barrier()
# local_data = {rank: dataloader.state_dict()}
# gathered_data = [None] * world_size
# dist.all_gather_object(gathered_data, local_data)
# if rank == 0:
# global_state_dict = {}
# for d in gathered_data:
# global_state_dict.update(d)
# torch.save(global_state_dict, "dataloader_state.pt")
assert torch.sum(masks[0]).item() == torch.count_nonzero(
embeddings[0]).item() // 4096
if args.vae_debug:
from diffusers.utils import export_to_video
from diffusers.video_processor import VideoProcessor
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.models.loader.component_loader import VAELoader
VAE_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/vae"
fastvideo_args = FastVideoArgs(
model_path=VAE_PATH,
vae_config=WanVAEConfig(load_encoder=False),
vae_precision="fp32")
fastvideo_args.device = device
vae_loader = VAELoader()
vae = vae_loader.load(model_path=VAE_PATH,
architecture="",
fastvideo_args=fastvideo_args)
videoprocessor = VideoProcessor(vae_scale_factor=8)
with torch.inference_mode():
video = vae.decode(latents[0].unsqueeze(0).to(device))
video = videoprocessor.postprocess_video(video)
video_path = os.path.join("/workspace/FastVideo/debug_videos",
infos["caption"][0][:50] + ".mp4")
export_to_video(video[0], video_path, fps=16)
# Move data to device
# latents = latents.to(device)
# embeddings = embeddings.to(device)
if world_size > 1:
dist.barrier()
# Benchmark
if rank == 0:
print(f"Benchmarking with batch_size={args.batch_size}")
start_time = time.time()
total_samples = 0
for i, (latents, embeddings, masks,
infos) in enumerate(tqdm.tqdm(dataloader, total=args.num_batches)):
if i >= args.num_batches:
break
# Move data to device
latents = latents.to(device)
embeddings = embeddings.to(device)
# Calculate actual batch size
batch_size = latents.size(0)
total_samples += batch_size
# Print progress only from rank 0
if rank == 0 and (i + 1) % 10 == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
print(
f"Batch {i+1}/{args.num_batches}, Speed: {samples_per_sec:.2f} samples/sec"
)
# Final statistics
if world_size > 1:
dist.barrier()
if rank == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
print("\nBenchmark Results:")
print(f"Total time: {elapsed:.2f} seconds")
print(f"Total samples: {total_samples}")
print(f"Average speed: {samples_per_sec:.2f} samples/sec")
print(f"Time per batch: {elapsed/args.num_batches*1000:.2f} ms")
if world_size > 1:
dist.destroy_process_group()
-349
View File
@@ -1,349 +0,0 @@
import json
import math
import os
import random
from collections import Counter
from os.path import join as opj
import numpy as np
import torch
import torchvision
from einops import rearrange
from PIL import Image
from torch.utils.data import Dataset
from fastvideo.utils.dataset_utils import DecordInit
from fastvideo.utils.logging_ import main_print
class SingletonMeta(type):
_instances: dict[type, 'SingletonMeta'] = {}
def __call__(cls, *args, **kwargs):
if cls not in cls._instances:
instance = super().__call__(*args, **kwargs)
cls._instances[cls] = instance
return cls._instances[cls]
class DataSetProg(metaclass=SingletonMeta):
def __init__(self) -> None:
self.cap_list: list[dict] = []
self.elements: list[int] = []
self.num_workers = 1
self.n_elements = 0
self.worker_elements: dict[int, list[int]] = {}
self.n_used_elements: dict[int, int] = {}
def set_cap_list(self, num_workers, cap_list, n_elements) -> None:
self.num_workers = num_workers
self.cap_list = cap_list
self.n_elements = n_elements
self.elements = list(range(n_elements))
random.shuffle(self.elements)
print(f"n_elements: {len(self.elements)}", flush=True)
for i in range(self.num_workers):
self.n_used_elements[i] = 0
per_worker = int(
math.ceil(len(self.elements) / float(self.num_workers)))
start = i * per_worker
end = min(start + per_worker, len(self.elements))
self.worker_elements[i] = self.elements[start:end]
def get_item(self, work_info) -> int:
worker_id = 0 if work_info is None else work_info.id
idx = self.worker_elements[worker_id][
self.n_used_elements[worker_id] %
len(self.worker_elements[worker_id])]
self.n_used_elements[worker_id] += 1
return idx
dataset_prog = DataSetProg()
def filter_resolution(h: int,
w: int,
max_h_div_w_ratio: float = 17 / 16,
min_h_div_w_ratio: float = 8 / 16) -> bool:
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
class T2V_dataset(Dataset):
def __init__(self,
args,
transform,
temporal_sample,
tokenizer,
transform_topcrop,
start_idx=0) -> None:
self.start_idx = start_idx
self.data = args.data_merge_path
self.num_frames = args.num_frames
self.train_fps = args.train_fps
self.use_image_num = args.use_image_num
self.transform = transform
self.transform_topcrop = transform_topcrop
self.temporal_sample = temporal_sample
self.tokenizer = tokenizer
self.text_max_length = args.text_max_length
self.cfg = args.cfg
self.speed_factor = args.speed_factor
self.max_height = args.max_height
self.max_width = args.max_width
self.drop_short_ratio = args.drop_short_ratio
assert self.speed_factor >= 1
self.v_decoder = DecordInit()
self.video_length_tolerance_range = args.video_length_tolerance_range
self.support_Chinese = True
if "mt5" not in args.text_encoder_name:
self.support_Chinese = False
cap_list = self.get_cap_list()
assert len(cap_list) > 0
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
self.lengths = self.sample_num_frames
n_elements = len(cap_list)
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
n_elements)
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
def set_checkpoint(self, n_used_elements):
for i in range(len(dataset_prog.n_used_elements)):
dataset_prog.n_used_elements[i] = n_used_elements
def __len__(self):
return dataset_prog.n_elements
def __getitem__(self, idx):
data = self.get_data(idx)
return data
def get_data(self, idx) -> dict:
path = dataset_prog.cap_list[idx]["path"]
if path.endswith(".mp4"):
return self.get_video(idx)
else:
return self.get_image(idx)
def get_video(self, idx) -> dict:
video_path = dataset_prog.cap_list[idx]["path"]
assert os.path.exists(video_path), f"file {video_path} do not exist!"
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
torchvision_video, _, metadata = torchvision.io.read_video(
video_path, output_format="TCHW")
video = torchvision_video[frame_indices]
video = self.transform(video)
video = rearrange(video, "t c h w -> c t h w")
video = video.to(torch.uint8)
assert video.dtype == torch.uint8
h, w = video.shape[-2:]
assert (
h / w <= 17 / 16 and h / w >= 8 / 16
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
video = video.float() / 127.5 - 1.0
text = dataset_prog.cap_list[idx]["cap"]
if not isinstance(text, list):
text = [text]
text = [random.choice(text)]
text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"]
cond_mask = text_tokens_and_mask["attention_mask"]
return dict(pixel_values=video,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=video_path,
fps=dataset_prog.cap_list[idx]["fps"],
duration=dataset_prog.cap_list[idx]["duration"])
def get_image(self, idx) -> dict:
image_data = dataset_prog.cap_list[
idx] # [{'path': path, 'cap': cap}, ...]
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
image = torch.from_numpy(np.array(image)) # [h, w, c]
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
# for i in image:
# h, w = i.shape[-2:]
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
image = (self.transform_topcrop(image) if "human_images"
in image_data["path"] else self.transform(image)
) # [1 C H W] -> num_img [1 C H W]
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
caps: list[str] = (image_data["cap"] if isinstance(
image_data["cap"], list) else [image_data["cap"]])
caps = [random.choice(caps)]
text = caps
input_ids, cond_mask = [], []
single_text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
single_text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"] # 1, l
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
return dict(
pixel_values=image,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=image_data["path"],
)
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
new_cap_list = []
sample_num_frames = []
cnt_too_long = 0
cnt_too_short = 0
cnt_no_cap = 0
cnt_no_resolution = 0
cnt_resolution_mismatch = 0
cnt_movie = 0
cnt_img = 0
for i in cap_list:
path = i["path"]
cap = i.get("cap", None)
# ======no caption=====
if cap is None:
cnt_no_cap += 1
continue
if path.endswith(".mp4"):
# ======no fps and duration=====
duration = i.get("duration", None)
fps = i.get("fps", None)
if fps is None or duration is None:
continue
# ======resolution mismatch=====
resolution = i.get("resolution", None)
if resolution is None:
cnt_no_resolution += 1
continue
else:
if (resolution.get("height", None) is None
or resolution.get("width", None) is None):
cnt_no_resolution += 1
continue
height, width = i["resolution"]["height"], i["resolution"][
"width"]
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
is_pick = filter_resolution(
height,
width,
max_h_div_w_ratio=hw_aspect_thr * aspect,
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
if not is_pick:
print("resolution mismatch")
cnt_resolution_mismatch += 1
continue
# import ipdb;ipdb.set_trace()
i["num_frames"] = math.ceil(fps * duration)
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
if i["num_frames"] / fps > self.video_length_tolerance_range * (
self.num_frames / self.train_fps * self.speed_factor
): # too long video is not suitable for this training stage (self.num_frames)
cnt_too_long += 1
continue
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
frame_interval = fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, i["num_frames"],
frame_interval).astype(int)
# comment out it to enable dynamic frames training
if (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio):
cnt_too_short += 1
continue
# too long video will be temporal-crop randomly
if len(frame_indices) > self.num_frames:
begin_index, end_index = self.temporal_sample(
len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
# frame_indices = frame_indices[:self.num_frames] # head crop
i["sample_frame_index"] = frame_indices.tolist()
new_cap_list.append(i)
i["sample_num_frames"] = len(
i["sample_frame_index"]
) # will use in dataloader(group sampler)
sample_num_frames.append(i["sample_num_frames"])
elif path.endswith(".jpg"): # image
cnt_img += 1
new_cap_list.append(i)
i["sample_num_frames"] = 1
sample_num_frames.append(i["sample_num_frames"])
else:
raise NameError(
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
)
# import ipdb;ipdb.set_trace()
main_print(
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
)
return new_cap_list, sample_num_frames
def decord_read(self, path, frame_indices) -> torch.Tensor:
decord_vr = self.v_decoder(path)
video_data = decord_vr.get_batch(frame_indices).asnumpy()
video_data = torch.from_numpy(video_data)
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
return video_data
def read_jsons(self, data) -> list[dict]:
cap_lists = []
with open(data) as f:
folder_anno = [
i.strip().split(",") for i in f.readlines()
if len(i.strip()) > 0
]
print(folder_anno)
for folder, anno in folder_anno:
with open(anno) as f:
sub_list = json.load(f)
for i in range(len(sub_list)):
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
cap_lists += sub_list
return cap_lists
def get_cap_list(self) -> list:
cap_lists = self.read_jsons(self.data)[self.start_idx:]
return cap_lists
-153
View File
@@ -1,153 +0,0 @@
import random
import torch
def _is_tensor_video_clip(clip) -> bool:
if not torch.is_tensor(clip):
raise TypeError(f"clip should be Tensor. Got {type(clip)}")
if not clip.ndimension() == 4:
raise ValueError(f"clip should be 4D. Got {clip.dim()}D")
return True
def crop(clip, i, j, h, w) -> torch.Tensor:
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
"""
if len(clip.size()) != 4:
raise ValueError("clip should be a 4D tensor")
return clip[..., i:i + h, j:j + w]
def resize(clip, target_size, interpolation_mode) -> torch.Tensor:
if len(target_size) != 2:
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
)
return torch.nn.functional.interpolate(
clip,
size=target_size,
mode=interpolation_mode,
align_corners=True,
antialias=True,
)
def center_crop_th_tw(clip, th, tw, top_crop) -> torch.Tensor:
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
# import ipdb;ipdb.set_trace()
h, w = clip.size(-2), clip.size(-1)
tr = th / tw
if h / w > tr:
new_h = int(w * tr)
new_w = w
else:
new_h = h
new_w = int(h / tr)
i = 0 if top_crop else int(round((h - new_h) / 2.0))
j = int(round((w - new_w) / 2.0))
return crop(clip, i, j, new_h, new_w)
def normalize_video(clip) -> torch.Tensor:
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
permute the dimensions of clip tensor
Args:
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
Return:
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
"""
_is_tensor_video_clip(clip)
if not clip.dtype == torch.uint8:
raise TypeError(
f"clip tensor should have data type uint8. Got {clip.dtype}")
# return clip.float().permute(3, 0, 1, 2) / 255.0
return clip.float() / 255.0
class CenterCropResizeVideo:
"""
First use the short side for cropping length,
center crop video, then resize to the specified size
"""
def __init__(
self,
size,
top_crop=False,
interpolation_mode="bilinear",
) -> None:
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
self.top_crop = top_crop
self.interpolation_mode = interpolation_mode
def __call__(self, clip) -> torch.Tensor:
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized / center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_center_crop = center_crop_th_tw(clip,
self.size[0],
self.size[1],
top_crop=self.top_crop)
clip_center_crop_resize = resize(
clip_center_crop,
target_size=self.size,
interpolation_mode=self.interpolation_mode,
)
return clip_center_crop_resize
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class Normalize255:
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
"""
def __init__(self) -> None:
pass
def __call__(self, clip) -> torch.Tensor:
"""
Args:
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
Return:
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
"""
return normalize_video(clip)
def __repr__(self) -> str:
return self.__class__.__name__
class TemporalRandomCrop:
"""Temporally crop the given frame indices at a random location.
Args:
size (int): Desired length of frames will be seen in the model.
"""
def __init__(self, size) -> None:
self.size = size
def __call__(self, total_frames) -> tuple[int, int]:
rand_end = max(0, total_frames - self.size - 1)
begin_index = random.randint(0, rand_end)
end_index = min(begin_index + self.size, total_frames)
return begin_index, end_index
-10
View File
@@ -1,10 +0,0 @@
from huggingface_hub import HfApi, upload_folder
api = HfApi()
repo_id = "weizhou03/HD-Mixkit-Finetune-Wan" # customize this
api.create_repo(repo_id=repo_id, repo_type="dataset")
upload_folder(repo_id=repo_id,
folder_path="/workspace/data/HD-Mixkit-Finetune-Wan",
repo_type="dataset",
path_in_repo="")
+1 -3
View File
@@ -5,8 +5,7 @@ from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size, get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size, get_world_group,
init_distributed_environment, initialize_model_parallel,
model_parallel_is_initialized)
init_distributed_environment, initialize_model_parallel)
from fastvideo.v1.distributed.utils import *
__all__ = [
@@ -18,5 +17,4 @@ __all__ = [
"get_tensor_model_parallel_world_size",
"cleanup_dist_env_and_memory",
"get_world_group",
"model_parallel_is_initialized",
]
@@ -1,182 +1,14 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/base_device_communicator.py
from typing import Any, Optional, Tuple
import torch
import torch.distributed as dist
from torch import Tensor
from torch.distributed import ProcessGroup, ReduceOp
class DistributedAutograd:
"""Collection of autograd functions for distributed operations.
This class provides custom autograd functions for distributed operations like all_reduce,
all_gather, and all_to_all. Each operation is implemented as a static inner class with
proper forward and backward implementations.
"""
class AllReduce(torch.autograd.Function):
"""Differentiable all_reduce operation.
The gradient of all_reduce is another all_reduce operation since the operation
combines values from all ranks equally.
"""
@staticmethod
def forward(ctx: Any,
group: ProcessGroup,
input_: Tensor,
op: Optional[dist.ReduceOp] = None) -> Tensor:
ctx.group = group
ctx.op = op
output = input_.clone()
dist.all_reduce(output, group=group, op=op)
return output
@staticmethod
def backward(ctx: Any,
grad_output: Tensor) -> Tuple[None, Tensor, None]:
grad_output = grad_output.clone()
dist.all_reduce(grad_output, group=ctx.group, op=ctx.op)
return None, grad_output, None
class AllGather(torch.autograd.Function):
"""Differentiable all_gather operation.
The operation gathers tensors from all ranks and concatenates them along a specified dimension.
The backward pass uses reduce_scatter to efficiently distribute gradients back to source ranks.
"""
@staticmethod
def forward(ctx: Any, group: ProcessGroup, input_: Tensor,
world_size: int, dim: int) -> Tensor:
ctx.group = group
ctx.world_size = world_size
ctx.dim = dim
ctx.input_shape = input_.shape
input_size = input_.size()
output_size = (input_size[0] * world_size, ) + input_size[1:]
output_tensor = torch.empty(output_size,
dtype=input_.dtype,
device=input_.device)
dist.all_gather_into_tensor(output_tensor, input_, group=group)
output_tensor = output_tensor.reshape((world_size, ) + input_size)
output_tensor = output_tensor.movedim(0, dim)
output_tensor = output_tensor.reshape(input_size[:dim] +
(world_size *
input_size[dim], ) +
input_size[dim + 1:])
return output_tensor
@staticmethod
def backward(ctx: Any,
grad_output: Tensor) -> Tuple[None, Tensor, None, None]:
# Split the gradient tensor along the gathered dimension
dim_size = grad_output.size(ctx.dim) // ctx.world_size
grad_chunks = grad_output.reshape(grad_output.shape[:ctx.dim] +
(ctx.world_size, dim_size) +
grad_output.shape[ctx.dim + 1:])
grad_chunks = grad_chunks.movedim(ctx.dim, 0)
# Each rank only needs its corresponding gradient
grad_input = torch.empty(ctx.input_shape,
dtype=grad_output.dtype,
device=grad_output.device)
dist.reduce_scatter_tensor(grad_input,
grad_chunks.contiguous(),
group=ctx.group)
return None, grad_input, None, None
class AllToAll4D(torch.autograd.Function):
"""Differentiable all_to_all operation specialized for 4D tensors.
This operation is particularly useful for attention operations where we need to
redistribute data across ranks for efficient parallel processing.
The operation supports two modes:
1. scatter_dim=2, gather_dim=1: Used for redistributing attention heads
2. scatter_dim=1, gather_dim=2: Used for redistributing sequence dimensions
"""
@staticmethod
def forward(ctx: Any, group: ProcessGroup, input_: Tensor,
world_size: int, scatter_dim: int,
gather_dim: int) -> Tensor:
ctx.group = group
ctx.world_size = world_size
ctx.scatter_dim = scatter_dim
ctx.gather_dim = gather_dim
if world_size == 1:
return input_
assert input_.dim(
) == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
if scatter_dim == 2 and gather_dim == 1:
bs, shard_seqlen, hc, hs = input_.shape
seqlen = shard_seqlen * world_size
shard_hc = hc // world_size
input_t = input_.reshape(bs, shard_seqlen, world_size, shard_hc,
hs).transpose(0, 2).contiguous()
output = torch.empty_like(input_t)
dist.all_to_all_single(output, input_t, group=group)
output = output.reshape(seqlen, bs, shard_hc,
hs).transpose(0, 1).contiguous()
output = output.reshape(bs, seqlen, shard_hc, hs)
return output
elif scatter_dim == 1 and gather_dim == 2:
bs, seqlen, shard_hc, hs = input_.shape
hc = shard_hc * world_size
shard_seqlen = seqlen // world_size
input_t = input_.reshape(bs, world_size, shard_seqlen, shard_hc,
hs)
input_t = input_t.transpose(0, 3).transpose(0, 1).contiguous()
input_t = input_t.reshape(world_size, shard_hc, shard_seqlen,
bs, hs)
output = torch.empty_like(input_t)
dist.all_to_all_single(output, input_t, group=group)
output = output.reshape(hc, shard_seqlen, bs, hs)
output = output.transpose(0, 2).contiguous()
output = output.reshape(bs, shard_seqlen, hc, hs)
return output
else:
raise RuntimeError(
f"Invalid scatter_dim={scatter_dim}, gather_dim={gather_dim}. "
f"Only (scatter_dim=2, gather_dim=1) and (scatter_dim=1, gather_dim=2) are supported."
)
@staticmethod
def backward(
ctx: Any,
grad_output: Tensor) -> Tuple[None, Tensor, None, None, None]:
if ctx.world_size == 1:
return None, grad_output, None, None, None
# For backward pass, we swap scatter_dim and gather_dim
output = DistributedAutograd.AllToAll4D.apply(
ctx.group, grad_output, ctx.world_size, ctx.gather_dim,
ctx.scatter_dim)
return None, output, None, None, None
from torch.distributed import ProcessGroup
class DeviceCommunicatorBase:
"""
Base class for device-specific communicator with autograd support.
Base class for device-specific communicator.
It can use the `cpu_group` to initialize the communicator.
If the device has PyTorch integration (PyTorch can recognize its
communication backend), the `device_group` will also be given.
@@ -184,8 +16,8 @@ class DeviceCommunicatorBase:
def __init__(self,
cpu_group: ProcessGroup,
device: Optional[torch.device] = None,
device_group: Optional[ProcessGroup] = None,
device: torch.device | None = None,
device_group: ProcessGroup | None = None,
unique_name: str = ""):
self.device = device or torch.device("cpu")
self.cpu_group = cpu_group
@@ -199,33 +31,40 @@ class DeviceCommunicatorBase:
self.rank_in_group = dist.get_group_rank(self.cpu_group,
self.global_rank)
def all_reduce(self,
input_: torch.Tensor,
op: Optional[dist.ReduceOp] = ReduceOp.SUM) -> torch.Tensor:
"""Performs an all_reduce operation with gradient support."""
return DistributedAutograd.AllReduce.apply(self.device_group, input_,
op)
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
dist.all_reduce(input_, group=self.device_group)
return input_
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
"""Performs an all_gather operation with gradient support."""
if dim < 0:
# Convert negative dim to positive.
dim += input_.dim()
return DistributedAutograd.AllGather.apply(self.device_group, input_,
self.world_size, dim)
def all_to_all_4D(self,
input_: torch.Tensor,
scatter_dim: int = 2,
gather_dim: int = 1) -> torch.Tensor:
"""Performs a 4D all-to-all operation with gradient support."""
return DistributedAutograd.AllToAll4D.apply(self.device_group, input_,
self.world_size,
scatter_dim, gather_dim)
input_size = input_.size()
# NOTE: we have to use concat-style all-gather here,
# stack-style all-gather has compatibility issues with
# torch.compile . see https://github.com/pytorch/pytorch/issues/138795
output_size = (input_size[0] * self.world_size, ) + input_size[1:]
# Allocate output tensor.
output_tensor = torch.empty(output_size,
dtype=input_.dtype,
device=input_.device)
# All-gather.
dist.all_gather_into_tensor(output_tensor,
input_,
group=self.device_group)
# Reshape
output_tensor = output_tensor.reshape((self.world_size, ) + input_size)
output_tensor = output_tensor.movedim(0, dim)
output_tensor = output_tensor.reshape(input_size[:dim] +
(self.world_size *
input_size[dim], ) +
input_size[dim + 1:])
return output_tensor
def gather(self,
input_: torch.Tensor,
dst: int = 0,
dim: int = -1) -> Optional[torch.Tensor]:
dim: int = -1) -> torch.Tensor | None:
"""
NOTE: We assume that the input tensor is on the same device across
all the ranks.
@@ -254,7 +93,82 @@ class DeviceCommunicatorBase:
output_tensor = None
return output_tensor
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
def all_to_all_4D(self,
input_: torch.Tensor,
scatter_dim: int = 2,
gather_dim: int = 1) -> torch.Tensor:
"""Specialized all-to-all operation for 4D tensors (e.g., for QKV matrices).
Args:
input_ (torch.Tensor): 4D input tensor to be scattered and gathered.
scatter_dim (int, optional): Dimension along which to scatter. Defaults to 2.
gather_dim (int, optional): Dimension along which to gather. Defaults to 1.
Returns:
torch.Tensor: Output tensor after all-to-all operation.
"""
# Bypass the function if we are using only 1 GPU.
if self.world_size == 1:
return input_
assert input_.dim(
) == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
if scatter_dim == 2 and gather_dim == 1:
# input: (bs, seqlen/P, hc, hs) output: (bs, seqlen, hc/P, hs)
bs, shard_seqlen, hc, hs = input_.shape
seqlen = shard_seqlen * self.world_size
shard_hc = hc // self.world_size
# Reshape and transpose for scattering
input_t = (input_.reshape(bs, shard_seqlen, self.world_size,
shard_hc, hs).transpose(0,
2).contiguous())
output = torch.empty_like(input_t)
torch.distributed.all_to_all_single(output,
input_t,
group=self.device_group)
torch.cuda.synchronize()
# Reshape and transpose back
output = output.reshape(seqlen, bs, shard_hc,
hs).transpose(0, 1).contiguous().reshape(
bs, seqlen, shard_hc, hs)
return output
elif scatter_dim == 1 and gather_dim == 2:
# input: (bs, seqlen, hc/P, hs) output: (bs, seqlen/P, hc, hs)
bs, seqlen, shard_hc, hs = input_.shape
hc = shard_hc * self.world_size
shard_seqlen = seqlen // self.world_size
# Reshape and transpose for scattering
input_t = (input_.reshape(bs, self.world_size, shard_seqlen,
shard_hc, hs).transpose(0, 3).transpose(
0, 1).contiguous().reshape(
self.world_size, shard_hc,
shard_seqlen, bs, hs))
output = torch.empty_like(input_t)
torch.distributed.all_to_all_single(output,
input_t,
group=self.device_group)
torch.cuda.synchronize()
# Reshape and transpose back
output = output.reshape(hc, shard_seqlen, bs,
hs).transpose(0, 2).contiguous().reshape(
bs, shard_seqlen, hc, hs)
return output
else:
raise RuntimeError(
"scatter_dim must be 1 or 2 and gather_dim must be 1 or 2")
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank."""
if dst is None:
@@ -264,7 +178,7 @@ class DeviceCommunicatorBase:
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: Optional[int] = None) -> torch.Tensor:
src: int | None = None) -> torch.Tensor:
"""Receives a tensor from the source rank."""
"""NOTE: `src` is the local rank of the source rank."""
if src is None:
@@ -1,8 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/cuda_communicator.py
from typing import Optional
import torch
from torch.distributed import ProcessGroup
@@ -14,37 +12,35 @@ class CudaCommunicator(DeviceCommunicatorBase):
def __init__(self,
cpu_group: ProcessGroup,
device: Optional[torch.device] = None,
device_group: Optional[ProcessGroup] = None,
device: torch.device | None = None,
device_group: ProcessGroup | None = None,
unique_name: str = ""):
super().__init__(cpu_group, device, device_group, unique_name)
from fastvideo.v1.distributed.device_communicators.pynccl import (
PyNcclCommunicator)
self.pynccl_comm: Optional[PyNcclCommunicator] = None
self.pynccl_comm: PyNcclCommunicator | None = None
if self.world_size > 1:
self.pynccl_comm = PyNcclCommunicator(
group=self.cpu_group,
device=self.device,
)
def all_reduce(self,
input_,
op: Optional[torch.distributed.ReduceOp] = None):
def all_reduce(self, input_):
pynccl_comm = self.pynccl_comm
assert pynccl_comm is not None
out = pynccl_comm.all_reduce(input_, op=op)
out = pynccl_comm.all_reduce(input_)
if out is None:
# fall back to the default all-reduce using PyTorch.
# this usually happens during testing.
# when we run the model, allreduce only happens for the TP
# group, where we always have either custom allreduce or pynccl.
out = input_.clone()
torch.distributed.all_reduce(out, group=self.device_group, op=op)
torch.distributed.all_reduce(out, group=self.device_group)
return out
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank."""
if dst is None:
@@ -59,7 +55,7 @@ class CudaCommunicator(DeviceCommunicatorBase):
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: Optional[int] = None) -> torch.Tensor:
src: int | None = None) -> torch.Tensor:
"""Receives a tensor from the source rank."""
"""NOTE: `src` is the local rank of the source rank."""
if src is None:
@@ -1,8 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/pynccl.py
from typing import Optional, Union
# ===================== import region =====================
import torch
import torch.distributed as dist
@@ -22,9 +20,9 @@ class PyNcclCommunicator:
def __init__(
self,
group: Union[ProcessGroup, StatelessProcessGroup],
device: Union[int, str, torch.device],
library_path: Optional[str] = None,
group: ProcessGroup | StatelessProcessGroup,
device: int | str | torch.device,
library_path: str | None = None,
):
"""
Args:
@@ -27,7 +27,7 @@
import ctypes
import platform
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
from typing import Any
import torch
from torch.distributed import ReduceOp
@@ -124,7 +124,7 @@ class ncclRedOpTypeEnum:
class Function:
name: str
restype: Any
argtypes: List[Any]
argtypes: list[Any]
class NCCLLibrary:
@@ -212,13 +212,13 @@ class NCCLLibrary:
# class attribute to store the mapping from the path to the library
# to avoid loading the same library multiple times
path_to_library_cache: Dict[str, Any] = {}
path_to_library_cache: dict[str, Any] = {}
# class attribute to store the mapping from library path
# to the corresponding dictionary
path_to_dict_mapping: Dict[str, Dict[str, Any]] = {}
path_to_dict_mapping: dict[str, dict[str, Any]] = {}
def __init__(self, so_file: Optional[str] = None):
def __init__(self, so_file: str | None = None):
so_file = so_file or find_nccl_library()
@@ -240,7 +240,7 @@ class NCCLLibrary:
raise e
if so_file not in NCCLLibrary.path_to_dict_mapping:
_funcs: Dict[str, Any] = {}
_funcs: dict[str, Any] = {}
for func in NCCLLibrary.exported_functions:
f = getattr(self.lib, func.name)
f.restype = func.restype
+50 -57
View File
@@ -27,15 +27,16 @@ import gc
import pickle
import weakref
from collections import namedtuple
from collections.abc import Callable
from contextlib import contextmanager
from dataclasses import dataclass
from multiprocessing import shared_memory
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
from typing import Any, Optional
from unittest.mock import patch
import torch
import torch.distributed
from torch.distributed import Backend, ProcessGroup, ReduceOp
from torch.distributed import Backend, ProcessGroup
import fastvideo.v1.envs as envs
from fastvideo.v1.distributed.device_communicators.base_device_communicator import (
@@ -57,15 +58,15 @@ TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
def _split_tensor_dict(
tensor_dict: Dict[str, Union[torch.Tensor, Any]]
) -> Tuple[List[Tuple[str, Any]], List[torch.Tensor]]:
tensor_dict: dict[str, torch.Tensor | Any]
) -> tuple[list[tuple[str, Any]], list[torch.Tensor]]:
"""Split the tensor dictionary into two parts:
1. A list of (key, value) pairs. If the value is a tensor, it is replaced
by its metadata.
2. A list of tensors.
"""
metadata_list: List[Tuple[str, Any]] = []
tensor_list: List[torch.Tensor] = []
metadata_list: list[tuple[str, Any]] = []
tensor_list: list[torch.Tensor] = []
for key, value in tensor_dict.items():
if isinstance(value, torch.Tensor):
# Note: we cannot use `value.device` here,
@@ -81,7 +82,7 @@ def _split_tensor_dict(
return metadata_list, tensor_list
_group_name_counter: Dict[str, int] = {}
_group_name_counter: dict[str, int] = {}
def _get_unique_name(name: str) -> str:
@@ -97,7 +98,7 @@ def _get_unique_name(name: str) -> str:
return newname
_groups: Dict[str, Callable[[], Optional["GroupCoordinator"]]] = {}
_groups: dict[str, Callable[[], Optional["GroupCoordinator"]]] = {}
def _register_group(group: "GroupCoordinator") -> None:
@@ -128,7 +129,7 @@ class GroupCoordinator:
# available attributes:
rank: int # global rank
ranks: List[int] # global ranks in the group
ranks: list[int] # global ranks in the group
world_size: int # size of the group
# difference between `local_rank` and `rank_in_group`:
# if we have a group of size 4 across two nodes:
@@ -143,16 +144,16 @@ class GroupCoordinator:
device_group: ProcessGroup # group for device communication
use_device_communicator: bool # whether to use device communicator
device_communicator: DeviceCommunicatorBase # device communicator
mq_broadcaster: Optional[Any] # shared memory broadcaster
mq_broadcaster: Any | None # shared memory broadcaster
def __init__(
self,
group_ranks: List[List[int]],
group_ranks: list[list[int]],
local_rank: int,
torch_distributed_backend: Union[str, Backend],
torch_distributed_backend: str | Backend,
use_device_communicator: bool,
use_message_queue_broadcaster: bool = False,
group_name: Optional[str] = None,
group_name: str | None = None,
):
group_name = group_name or "anonymous"
self.unique_name = _get_unique_name(group_name)
@@ -243,8 +244,8 @@ class GroupCoordinator:
return self.ranks[(rank_in_group - 1) % world_size]
@contextmanager
def graph_capture(
self, graph_capture_context: Optional[GraphCaptureContext] = None):
def graph_capture(self,
graph_capture_context: GraphCaptureContext | None = None):
if graph_capture_context is None:
stream = torch.cuda.Stream()
graph_capture_context = GraphCaptureContext(stream)
@@ -260,11 +261,7 @@ class GroupCoordinator:
with torch.cuda.stream(stream):
yield graph_capture_context
def all_reduce(
self,
input_: torch.Tensor,
op: Optional[torch.distributed.ReduceOp] = ReduceOp.SUM
) -> torch.Tensor:
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
"""
User-facing all-reduce function before we actually call the
all-reduce operation.
@@ -287,14 +284,10 @@ class GroupCoordinator:
return torch.ops.vllm.all_reduce(input_,
group_name=self.unique_name)
else:
return self._all_reduce_out_place(input_, op=op)
return self._all_reduce_out_place(input_)
def _all_reduce_out_place(
self,
input_: torch.Tensor,
op: Optional[torch.distributed.ReduceOp] = ReduceOp.SUM
) -> torch.Tensor:
return self.device_communicator.all_reduce(input_, op=op)
def _all_reduce_out_place(self, input_: torch.Tensor) -> torch.Tensor:
return self.device_communicator.all_reduce(input_)
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
world_size = self.world_size
@@ -309,7 +302,7 @@ class GroupCoordinator:
def gather(self,
input_: torch.Tensor,
dst: int = 0,
dim: int = -1) -> Optional[torch.Tensor]:
dim: int = -1) -> torch.Tensor | None:
"""
NOTE: We assume that the input tensor is on the same device across
all the ranks.
@@ -345,7 +338,7 @@ class GroupCoordinator:
group=self.device_group)
return input_
def broadcast_object(self, obj: Optional[Any] = None, src: int = 0):
def broadcast_object(self, obj: Any | None = None, src: int = 0):
"""Broadcast the input object.
NOTE: `src` is the local rank of the source rank.
"""
@@ -370,9 +363,9 @@ class GroupCoordinator:
return recv[0]
def broadcast_object_list(self,
obj_list: List[Any],
obj_list: list[Any],
src: int = 0,
group: Optional[ProcessGroup] = None):
group: ProcessGroup | None = None):
"""Broadcast the input object list.
NOTE: `src` is the local rank of the source rank.
"""
@@ -452,11 +445,11 @@ class GroupCoordinator:
def broadcast_tensor_dict(
self,
tensor_dict: Optional[Dict[str, Union[torch.Tensor, Any]]] = None,
tensor_dict: dict[str, torch.Tensor | Any] | None = None,
src: int = 0,
group: Optional[ProcessGroup] = None,
metadata_group: Optional[ProcessGroup] = None
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
group: ProcessGroup | None = None,
metadata_group: ProcessGroup | None = None
) -> dict[str, torch.Tensor | Any] | None:
"""Broadcast the input tensor dictionary.
NOTE: `src` is the local rank of the source rank.
"""
@@ -470,7 +463,7 @@ class GroupCoordinator:
rank_in_group = self.rank_in_group
if rank_in_group == src:
metadata_list: List[Tuple[Any, Any]] = []
metadata_list: list[tuple[Any, Any]] = []
assert isinstance(
tensor_dict,
dict), (f"Expecting a dictionary, got {type(tensor_dict)}")
@@ -537,10 +530,10 @@ class GroupCoordinator:
def send_tensor_dict(
self,
tensor_dict: Dict[str, Union[torch.Tensor, Any]],
dst: Optional[int] = None,
tensor_dict: dict[str, torch.Tensor | Any],
dst: int | None = None,
all_gather_group: Optional["GroupCoordinator"] = None,
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
) -> dict[str, torch.Tensor | Any] | None:
"""Send the input tensor dictionary.
NOTE: `dst` is the local rank of the source rank.
"""
@@ -560,7 +553,7 @@ class GroupCoordinator:
dst = (self.rank_in_group + 1) % self.world_size
assert dst < self.world_size, f"Invalid dst rank ({dst})"
metadata_list: List[Tuple[Any, Any]] = []
metadata_list: list[tuple[Any, Any]] = []
assert isinstance(
tensor_dict,
dict), f"Expecting a dictionary, got {type(tensor_dict)}"
@@ -591,9 +584,9 @@ class GroupCoordinator:
def recv_tensor_dict(
self,
src: Optional[int] = None,
src: int | None = None,
all_gather_group: Optional["GroupCoordinator"] = None,
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
) -> dict[str, torch.Tensor | Any] | None:
"""Recv the input tensor dictionary.
NOTE: `src` is the local rank of the source rank.
"""
@@ -614,7 +607,7 @@ class GroupCoordinator:
assert src < self.world_size, f"Invalid src rank ({src})"
recv_metadata_list = self.recv_object(src=src)
tensor_dict: Dict[str, Any] = {}
tensor_dict: dict[str, Any] = {}
for key, value in recv_metadata_list:
if isinstance(value, TensorMetadata):
tensor = torch.empty(value.size,
@@ -664,7 +657,7 @@ class GroupCoordinator:
"""
torch.distributed.barrier(group=self.cpu_group)
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank."""
self.device_communicator.send(tensor, dst)
@@ -672,7 +665,7 @@ class GroupCoordinator:
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: Optional[int] = None) -> torch.Tensor:
src: int | None = None) -> torch.Tensor:
"""Receives a tensor from the source rank."""
"""NOTE: `src` is the local rank of the source rank."""
return self.device_communicator.recv(size, dtype, src)
@@ -690,7 +683,7 @@ class GroupCoordinator:
self.mq_broadcaster = None
_WORLD: Optional[GroupCoordinator] = None
_WORLD: GroupCoordinator | None = None
def get_world_group() -> GroupCoordinator:
@@ -698,7 +691,7 @@ def get_world_group() -> GroupCoordinator:
return _WORLD
def init_world_group(ranks: List[int], local_rank: int,
def init_world_group(ranks: list[int], local_rank: int,
backend: str) -> GroupCoordinator:
return GroupCoordinator(
group_ranks=[ranks],
@@ -710,11 +703,11 @@ def init_world_group(ranks: List[int], local_rank: int,
def init_model_parallel_group(
group_ranks: List[List[int]],
group_ranks: list[list[int]],
local_rank: int,
backend: str,
use_message_queue_broadcaster: bool = False,
group_name: Optional[str] = None,
group_name: str | None = None,
) -> GroupCoordinator:
return GroupCoordinator(
@@ -727,7 +720,7 @@ def init_model_parallel_group(
)
_TP: Optional[GroupCoordinator] = None
_TP: GroupCoordinator | None = None
def get_tp_group() -> GroupCoordinator:
@@ -786,7 +779,7 @@ def init_distributed_environment(
"world group already initialized with a different world size")
_SP: Optional[GroupCoordinator] = None
_SP: GroupCoordinator | None = None
def get_sp_group() -> GroupCoordinator:
@@ -797,7 +790,7 @@ def get_sp_group() -> GroupCoordinator:
def initialize_model_parallel(
tensor_model_parallel_size: int = 1,
sequence_model_parallel_size: int = 1,
backend: Optional[str] = None,
backend: str | None = None,
) -> None:
"""
Initialize model parallel groups.
@@ -866,7 +859,7 @@ def get_sequence_model_parallel_rank() -> int:
def ensure_model_parallel_initialized(
tensor_model_parallel_size: int,
sequence_model_parallel_size: int,
backend: Optional[str] = None,
backend: str | None = None,
) -> None:
"""Helper to initialize model parallel groups if they are not initialized,
or ensure tensor-parallel, sequence-parallel sizes
@@ -977,8 +970,8 @@ 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 in_the_same_node_as(pg: ProcessGroup | StatelessProcessGroup,
source_rank: int = 0) -> list[bool]:
"""
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
@@ -1064,7 +1057,7 @@ def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
def initialize_tensor_parallel_group(
tensor_model_parallel_size: int = 1,
backend: Optional[str] = None,
backend: str | None = None,
group_name_suffix: str = "") -> GroupCoordinator:
"""Initialize a tensor parallel group for a specific model.
@@ -1128,7 +1121,7 @@ def initialize_tensor_parallel_group(
def initialize_sequence_parallel_group(
sequence_model_parallel_size: int = 1,
backend: Optional[str] = None,
backend: str | None = None,
group_name_suffix: str = "") -> GroupCoordinator:
"""Initialize a sequence parallel group for a specific model.
+7 -6
View File
@@ -9,7 +9,8 @@ import dataclasses
import pickle
import time
from collections import deque
from typing import Any, Deque, Dict, Optional, Sequence, Tuple
from collections.abc import Sequence
from typing import Any
import torch
from torch.distributed import TCPStore
@@ -72,15 +73,15 @@ class StatelessProcessGroup:
data_expiration_seconds: int = 3600 # 1 hour
# dst rank -> counter
send_dst_counter: Dict[int, int] = dataclasses.field(default_factory=dict)
send_dst_counter: dict[int, int] = dataclasses.field(default_factory=dict)
# src rank -> counter
recv_src_counter: Dict[int, int] = dataclasses.field(default_factory=dict)
recv_src_counter: dict[int, int] = dataclasses.field(default_factory=dict)
broadcast_send_counter: int = 0
broadcast_recv_src_counter: Dict[int, int] = dataclasses.field(
broadcast_recv_src_counter: dict[int, int] = dataclasses.field(
default_factory=dict)
# A deque to store the data entries, with key and timestamp.
entries: Deque[Tuple[str, float]] = dataclasses.field(default_factory=deque)
entries: deque[tuple[str, float]] = dataclasses.field(default_factory=deque)
def __post_init__(self):
assert self.rank < self.world_size
@@ -114,7 +115,7 @@ class StatelessProcessGroup:
self.recv_src_counter[src] += 1
return obj
def broadcast_obj(self, obj: Optional[Any], src: int) -> Any:
def broadcast_obj(self, obj: Any | None, src: int) -> Any:
"""Broadcast an object from a source rank to all other ranks.
It does not clean up after all ranks have received the object.
Use it for limited times, e.g., for initialization.
+6 -6
View File
@@ -4,7 +4,7 @@
import argparse
import dataclasses
import os
from typing import Any, Dict, List, Optional, cast
from typing import Any, cast
from fastvideo import PipelineConfig, VideoGenerator
from fastvideo.v1.configs.sample.base import SamplingParam
@@ -26,11 +26,11 @@ class GenerateSubcommand(CLISubcommand):
self.init_arg_names = self._get_init_arg_names()
self.generation_arg_names = self._get_generation_arg_names()
def _get_init_arg_names(self) -> List[str]:
def _get_init_arg_names(self) -> list[str]:
"""Get names of arguments for VideoGenerator initialization"""
return ["num_gpus", "tp_size", "sp_size", "model_path"]
def _get_generation_arg_names(self) -> List[str]:
def _get_generation_arg_names(self) -> list[str]:
"""Get names of arguments for generate_video method"""
return [field.name for field in dataclasses.fields(SamplingParam)]
@@ -130,13 +130,13 @@ class GenerateSubcommand(CLISubcommand):
return cast(FlexibleArgumentParser, generate_parser)
def cmd_init() -> List[CLISubcommand]:
def cmd_init() -> list[CLISubcommand]:
return [GenerateSubcommand()]
def update_config_from_args(config: Any,
args_dict: Dict[str, Any],
prefix: Optional[str] = None) -> None:
args_dict: dict[str, Any],
prefix: str | None = None) -> None:
"""
Update configuration object from arguments dictionary.
+1 -3
View File
@@ -1,14 +1,12 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/main.py
from typing import List
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.v1.entrypoints.cli.generate import cmd_init as generate_cmd_init
from fastvideo.v1.utils import FlexibleArgumentParser
def cmd_init() -> List[CLISubcommand]:
def cmd_init() -> list[CLISubcommand]:
"""Initialize all commands from separate modules"""
commands = []
commands.extend(generate_cmd_init())
+2 -3
View File
@@ -4,7 +4,6 @@ import argparse
import os
import subprocess
import sys
from typing import List, Optional
from fastvideo.v1.logger import init_logger
@@ -19,8 +18,8 @@ class RaiseNotImplementedAction(argparse.Action):
def launch_distributed(num_gpus: int,
args: List[str],
master_port: Optional[int] = None) -> int:
args: list[str],
master_port: int | None = None) -> int:
"""
Launch a distributed job with the given arguments
@@ -1,29 +0,0 @@
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from fastvideo.v1.pipelines.wan.wan_latent_pipeline import WanLatentPipeline
def main():
print("Starting data preprocessor")
pipeline = WanLatentPipeline.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
train_dataset = getdataset(args)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
for batch in train_dataloader:
pipeline(batch)
if __name__ == "__main__":
main()
+6 -8
View File
@@ -10,7 +10,7 @@ import gc
import math
import os
import time
from typing import Any, Dict, List, Optional, Union
from typing import Any
import imageio
import numpy as np
@@ -53,11 +53,9 @@ class VideoGenerator:
@classmethod
def from_pretrained(cls,
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[
Union[str
| PipelineConfig]] = None,
device: str | None = None,
torch_dtype: torch.dtype | None = None,
pipeline_config: str | PipelineConfig | None = None,
**kwargs) -> "VideoGenerator":
"""
Create a video generator from a pretrained model.
@@ -128,9 +126,9 @@ class VideoGenerator:
def generate_video(
self,
prompt: str,
sampling_param: Optional[SamplingParam] = None,
sampling_param: SamplingParam | None = None,
**kwargs,
) -> Union[Dict[str, Any], List[np.ndarray]]:
) -> dict[str, Any] | list[np.ndarray]:
"""
Generate a video based on the given prompt.
+13 -12
View File
@@ -2,28 +2,29 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/envs.py
import os
from typing import TYPE_CHECKING, Any, Callable, Dict, Optional
from collections.abc import Callable
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60
FASTVIDEO_NCCL_SO_PATH: Optional[str] = None
LD_LIBRARY_PATH: Optional[str] = None
FASTVIDEO_NCCL_SO_PATH: str | None = None
LD_LIBRARY_PATH: str | None = None
LOCAL_RANK: int = 0
CUDA_VISIBLE_DEVICES: Optional[str] = None
CUDA_VISIBLE_DEVICES: str | None = None
FASTVIDEO_CACHE_ROOT: str = os.path.expanduser("~/.cache/fastvideo")
FASTVIDEO_CONFIG_ROOT: str = os.path.expanduser("~/.config/fastvideo")
FASTVIDEO_CONFIGURE_LOGGING: int = 1
FASTVIDEO_LOGGING_LEVEL: str = "INFO"
FASTVIDEO_LOGGING_PREFIX: str = ""
FASTVIDEO_LOGGING_CONFIG_PATH: Optional[str] = None
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: Optional[str] = None
FASTVIDEO_ATTENTION_CONFIG: Optional[str] = None
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_ATTENTION_CONFIG: str | None = None
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: Optional[str] = None
NVCC_THREADS: Optional[str] = None
CMAKE_BUILD_TYPE: Optional[str] = None
MAX_JOBS: str | None = None
NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: str | None = None
VERBOSE: bool = False
FASTVIDEO_SERVER_DEV_MODE: bool = False
@@ -42,7 +43,7 @@ def get_default_config_root() -> str:
)
def maybe_convert_int(value: Optional[str]) -> Optional[int]:
def maybe_convert_int(value: str | None) -> int | None:
if value is None:
return None
return int(value)
@@ -53,7 +54,7 @@ def maybe_convert_int(value: Optional[str]) -> Optional[int]:
# begin-env-vars-definition
environment_variables: Dict[str, Callable[[], Any]] = {
environment_variables: dict[str, Callable[[], Any]] = {
# ================== Installation Time Env Vars ==================
+18 -359
View File
@@ -4,9 +4,10 @@
import argparse
import dataclasses
from collections.abc import Callable
from contextlib import contextmanager
from dataclasses import field
from typing import Any, Callable, List, Optional, Tuple
from typing import Any
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.logger import init_logger
@@ -38,17 +39,17 @@ class FastVideoArgs:
# HuggingFace specific parameters
trust_remote_code: bool = False
revision: Optional[str] = None
revision: str | None = None
# Parallelism
num_gpus: int = 1
tp_size: Optional[int] = None
sp_size: Optional[int] = None
dist_timeout: Optional[int] = None # timeout for torch.distributed
tp_size: int | None = None
sp_size: int | None = None
dist_timeout: int | None = None # timeout for torch.distributed
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
flow_shift: float | None = None
output_type: str = "pil"
@@ -70,40 +71,36 @@ class FastVideoArgs:
# Text encoder configuration
DEFAULT_TEXT_ENCODER_PRECISIONS = (
"fp16",
# "fp16",
"fp16",
)
text_encoder_precisions: Tuple[str, ...] = field(
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), ))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: Tuple[Callable[[Any], Any], ...] = field(
postprocess_text_funcs: tuple[Callable[[Any], Any], ...] = field(
default_factory=lambda: (postprocess_text, ))
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
mask_strategy_file_path: str | None = None
enable_torch_compile: bool = False
use_cpu_offload: bool = False
disable_autocast: bool = False
# StepVideo specific parameters
pos_magic: Optional[str] = None
neg_magic: Optional[str] = None
timesteps_scale: Optional[bool] = None
pos_magic: str | None = None
neg_magic: str | None = None
timesteps_scale: bool | None = None
# Logging
log_level: str = "info"
# Inference parameters
device_str: Optional[str] = None
device_str: str | None = None
device = None
@property
def training_mode(self) -> bool:
return not self.inference_mode
def __post_init__(self):
pass
@@ -136,13 +133,6 @@ class FastVideoArgs:
help="The distributed executor backend to use",
)
parser.add_argument(
"--inference-mode",
action=StoreBoolean,
default=FastVideoArgs.inference_mode,
help="Whether to use inference mode",
)
# HuggingFace specific parameters
parser.add_argument(
"--trust-remote-code",
@@ -387,7 +377,7 @@ class FastVideoArgs:
_current_fastvideo_args = None
def prepare_fastvideo_args(argv: List[str]) -> FastVideoArgs:
def prepare_fastvideo_args(argv: list[str]) -> FastVideoArgs:
"""
Prepare the inference arguments from the command line arguments.
@@ -434,334 +424,3 @@ def get_current_fastvideo_args() -> FastVideoArgs:
# TODO(will): may need to handle this for CI.
raise ValueError("Current fastvideo args is not set.")
return _current_fastvideo_args
@dataclasses.dataclass
class TrainingArgs(FastVideoArgs):
"""
Training arguments. Inherits from FastVideoArgs and adds training-specific
arguments. If there are any conflicts, the training arguments will take
precedence.
"""
data_path: str = ""
dataloader_num_workers: int = 0
num_height: int = 0
num_width: int = 0
num_frames: int = 0
train_batch_size: int = 0
num_latent_t: int = 0
group_frame: bool = False
group_resolution: bool = False
# text encoder & vae & diffusion model
pretrained_model_name_or_path: str = ""
dit_model_name_or_path: str = ""
cache_dir: str = ""
# diffusion setting
ema_decay: float = 0.0
ema_start_step: int = 0
cfg: float = 0.0
precondition_outputs: bool = False
# validation & logs
validation_prompt_dir: str = ""
validation_sampling_steps: str = ""
validation_guidance_scale: str = ""
validation_steps: float = 0.0
log_validation: bool = False
tracker_project_name: str = ""
# seed: int
# output
output_dir: str = ""
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: str = ""
logging_dir: str = ""
# optimizer & scheduler
num_train_epochs: int = 0
max_train_steps: int = 0
gradient_accumulation_steps: int = 0
learning_rate: float = 0.0
scale_lr: bool = False
lr_scheduler: str = ""
lr_warmup_steps: int = 0
max_grad_norm: float = 0.0
gradient_checkpointing: bool = False
selective_checkpointing: float = 0.0
allow_tf32: bool = False
mixed_precision: str = ""
train_sp_batch_size: int = 0
fsdp_sharding_startegy: str = ""
weighting_scheme: str = ""
logit_mean: float = 0.0
logit_std: float = 1.0
mode_scale: float = 0.0
num_euler_timesteps: int = 0
lr_num_cycles: int = 0
lr_power: float = 0.0
not_apply_cfg_solver: bool = False
distill_cfg: float = 0.0
scheduler_type: str = ""
linear_quadratic_threshold: float = 0.0
linear_range: float = 0.0
weight_decay: float = 0.0
use_ema: bool = False
multi_phased_distill_schedule: str = ""
pred_decay_weight: float = 0.0
pred_decay_type: str = ""
hunyuan_teacher_disable_cfg: bool = False
# master_weight_type
master_weight_type: str = ""
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
# Get all fields from the dataclass
attrs = [attr.name for attr in dataclasses.fields(cls)]
# Create a dictionary of attribute values, with defaults for missing attributes
kwargs = {}
for attr in attrs:
# Handle renamed attributes or those with multiple CLI names
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
# Use getattr with default value from the dataclass for potentially missing attributes
else:
default_value = getattr(cls, attr, None)
kwargs[attr] = getattr(args, attr, default_value)
return cls(**kwargs)
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
parser.add_argument("--data-path",
type=str,
required=True,
help="Path to parquet files")
parser.add_argument("--dataloader-num-workers",
type=int,
required=True,
help="Number of workers for dataloader")
parser.add_argument("--num-height",
type=int,
required=True,
help="Number of heights")
parser.add_argument("--num-width",
type=int,
required=True,
help="Number of widths")
parser.add_argument("--num-frames",
type=int,
required=True,
help="Number of frames")
# Training batch and model configuration
parser.add_argument("--train-batch-size",
type=int,
required=True,
help="Training batch size")
parser.add_argument("--num-latent-t",
type=int,
required=True,
help="Number of latent time steps")
parser.add_argument("--group-frame",
action=StoreBoolean,
help="Whether to group frames during training")
parser.add_argument("--group-resolution",
action=StoreBoolean,
help="Whether to group resolutions during training")
# Model paths
parser.add_argument("--pretrained-model-name-or-path",
type=str,
required=True,
help="Path to pretrained model or model name")
parser.add_argument("--dit-model-name-or-path",
type=str,
required=False,
help="Path to DiT model or model name")
parser.add_argument("--cache-dir",
type=str,
help="Directory to cache models")
# Diffusion settings
parser.add_argument("--ema-decay",
type=float,
default=0.999,
help="EMA decay rate")
parser.add_argument("--ema-start-step",
type=int,
default=0,
help="Step to start EMA")
parser.add_argument("--cfg",
type=float,
help="Classifier-free guidance scale")
parser.add_argument(
"--precondition-outputs",
action=StoreBoolean,
help="Whether to precondition the outputs of the model")
# Validation and logging
parser.add_argument("--validation-prompt-dir",
type=str,
help="Directory containing validation prompts")
parser.add_argument("--validation-sampling-steps",
type=str,
help="Validation sampling steps")
parser.add_argument("--validation-guidance-scale",
type=str,
help="Validation guidance scale")
parser.add_argument("--validation-steps",
type=float,
help="Number of validation steps")
parser.add_argument("--log-validation",
action=StoreBoolean,
help="Whether to log validation results")
parser.add_argument("--tracker-project-name",
type=str,
help="Project name for tracking")
# Output configuration
parser.add_argument("--output-dir",
type=str,
required=True,
help="Output directory for checkpoints and logs")
parser.add_argument("--checkpoints-total-limit",
type=int,
help="Maximum number of checkpoints to keep")
parser.add_argument("--checkpointing-steps",
type=int,
help="Steps between checkpoints")
parser.add_argument("--resume-from-checkpoint",
type=str,
help="Path to checkpoint to resume from")
parser.add_argument("--logging-dir",
type=str,
help="Directory for logging")
# Training configuration
parser.add_argument("--num-train-epochs",
type=int,
help="Number of training epochs")
parser.add_argument("--max-train-steps",
type=int,
help="Maximum number of training steps")
parser.add_argument("--gradient-accumulation-steps",
type=int,
help="Number of steps to accumulate gradients")
parser.add_argument("--learning-rate",
type=float,
required=True,
help="Learning rate")
parser.add_argument("--scale-lr",
action=StoreBoolean,
help="Whether to scale learning rate")
parser.add_argument("--lr-scheduler",
type=str,
default="constant",
help="Learning rate scheduler type")
parser.add_argument("--lr-warmup-steps",
type=int,
default=10,
help="Number of warmup steps for learning rate")
parser.add_argument("--max-grad-norm",
type=float,
help="Maximum gradient norm")
parser.add_argument("--gradient-checkpointing",
action=StoreBoolean,
help="Whether to use gradient checkpointing")
parser.add_argument("--selective-checkpointing",
type=float,
help="Selective checkpointing threshold")
parser.add_argument("--allow-tf32",
action=StoreBoolean,
help="Whether to allow TF32")
parser.add_argument("--mixed-precision",
type=str,
help="Mixed precision training type")
parser.add_argument("--train-sp-batch-size",
type=int,
help="Training spatial parallelism batch size")
parser.add_argument("--fsdp-sharding-strategy",
type=str,
help="FSDP sharding strategy")
parser.add_argument(
"--weighting_scheme",
type=str,
default="uniform",
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "uniform"],
)
parser.add_argument(
"--logit_mean",
type=float,
default=0.0,
help="mean to use when using the `'logit_normal'` weighting scheme.",
)
parser.add_argument(
"--logit_std",
type=float,
default=1.0,
help="std to use when using the `'logit_normal'` weighting scheme.",
)
parser.add_argument(
"--mode_scale",
type=float,
default=1.29,
help=
"Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
)
# Additional training parameters
parser.add_argument("--num-euler-timesteps",
type=int,
help="Number of Euler timesteps")
parser.add_argument("--lr-num-cycles",
type=int,
help="Number of learning rate cycles")
parser.add_argument("--lr-power",
type=float,
help="Learning rate power")
parser.add_argument("--not-apply-cfg-solver",
action=StoreBoolean,
help="Whether to not apply CFG solver")
parser.add_argument("--distill-cfg",
type=float,
help="Distillation CFG scale")
parser.add_argument("--scheduler-type", type=str, help="Scheduler type")
parser.add_argument("--linear-quadratic-threshold",
type=float,
help="Linear quadratic threshold")
parser.add_argument("--linear-range", type=float, help="Linear range")
parser.add_argument("--weight-decay", type=float, help="Weight decay")
parser.add_argument("--use-ema",
action=StoreBoolean,
help="Whether to use EMA")
parser.add_argument("--multi-phased-distill-schedule",
type=str,
help="Multi-phased distillation schedule")
parser.add_argument("--pred-decay-weight",
type=float,
help="Prediction decay weight")
parser.add_argument("--pred-decay-type",
type=str,
help="Prediction decay type")
parser.add_argument("--hunyuan-teacher-disable-cfg",
action=StoreBoolean,
help="Whether to disable CFG for Hunyuan teacher")
parser.add_argument("--master-weight-type",
type=str,
help="Master weight type")
return parser
+5 -5
View File
@@ -5,7 +5,7 @@ import time
from collections import defaultdict
from contextlib import contextmanager
from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING
import torch
@@ -37,10 +37,10 @@ class ForwardContext:
# attn_layers: Dict[str, Any]
# TODO: extend to support per-layer dynamic forward context
attn_metadata: "AttentionMetadata" # set dynamically for each forward pass
forward_batch: Optional[ForwardBatch] = None
forward_batch: ForwardBatch | None = None
_forward_context: Optional[ForwardContext] = None
_forward_context: ForwardContext | None = None
def get_forward_context() -> ForwardContext:
@@ -55,8 +55,8 @@ def get_forward_context() -> ForwardContext:
@contextmanager
def set_forward_context(current_timestep,
attn_metadata,
forward_batch: Optional[ForwardBatch] = None,
fastvideo_args: Optional[FastVideoArgs] = None):
forward_batch: ForwardBatch | None = None,
fastvideo_args: FastVideoArgs | None = None):
"""A context manager that stores the current forward context,
can be attention metadata, etc.
Here we can inject common logic for every model forward pass.
+3 -3
View File
@@ -8,7 +8,7 @@ This module provides classes and functions for running inference with diffusion
"""
import time
from typing import Any, Dict
from typing import Any
import torch
@@ -83,7 +83,7 @@ class InferenceEngine:
self,
prompt: str,
fastvideo_args: FastVideoArgs,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
Run inference with the pipeline.
@@ -96,7 +96,7 @@ class InferenceEngine:
Returns:
A dictionary containing the generated videos and metadata.
"""
out_dict: Dict[str, Any] = dict()
out_dict: dict[str, Any] = dict()
num_videos_per_prompt = fastvideo_args.num_videos
seed = fastvideo_args.seed
+3 -2
View File
@@ -1,7 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/custom_op.py
from typing import Any, Callable, Dict, Type
from collections.abc import Callable
from typing import Any
import torch.nn as nn
@@ -81,7 +82,7 @@ class CustomOp(nn.Module):
# Examples:
# - MyOp.enabled()
# - op_registry["my_op"].enabled()
op_registry: Dict[str, Type['CustomOp']] = {}
op_registry: dict[str, type['CustomOp']] = {}
# Decorator to register custom ops.
@classmethod
+4 -5
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/layernorm.py
"""Custom normalization layers."""
from typing import Optional, Tuple, Union
import torch
import torch.nn as nn
@@ -22,7 +21,7 @@ class RMSNorm(CustomOp):
hidden_size: int,
eps: float = 1e-6,
dtype: torch.dtype = torch.float32,
var_hidden_size: Optional[int] = None,
var_hidden_size: int | None = None,
has_weight: bool = True,
) -> None:
super().__init__()
@@ -40,8 +39,8 @@ class RMSNorm(CustomOp):
def forward_native(
self,
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
residual: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""PyTorch-native implementation equivalent to forward()."""
orig_dtype = x.dtype
x = x.to(torch.float32)
@@ -130,7 +129,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor, shift: torch.Tensor,
scale: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
scale: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
Apply gated residual connection, followed by layernorm and
scale/shift in a single fused operation.
+34 -42
View File
@@ -2,7 +2,6 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/linear.py
from abc import abstractmethod
from typing import Optional, Union
import torch
import torch.nn.functional as F
@@ -40,7 +39,7 @@ WEIGHT_LOADER_V2_SUPPORTED = [
def adjust_scalar_to_fused_array(
param: torch.Tensor, loaded_weight: torch.Tensor,
shard_id: Union[str, int]) -> tuple[torch.Tensor, torch.Tensor]:
shard_id: str | int) -> tuple[torch.Tensor, torch.Tensor]:
"""For fused modules (QKV and MLP) we have an array of length
N that holds 1 scale for each "logical" matrix. So the param
is an array of length N. The loaded_weight corresponds to
@@ -91,7 +90,7 @@ class LinearMethodBase(QuantizeMethodBase):
def apply(self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None) -> torch.Tensor:
bias: torch.Tensor | None = None) -> torch.Tensor:
"""Apply the weights in layer to the input tensor.
Expects create_weights to have been called before on the layer."""
raise NotImplementedError
@@ -116,7 +115,7 @@ class UnquantizedLinearMethod(LinearMethodBase):
def apply(self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None) -> torch.Tensor:
bias: torch.Tensor | None = None) -> torch.Tensor:
return F.linear(x, layer.weight, bias)
@@ -138,8 +137,8 @@ class LinearBase(torch.nn.Module):
input_size: int,
output_size: int,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__()
@@ -152,14 +151,13 @@ class LinearBase(torch.nn.Module):
params_dtype = torch.get_default_dtype()
self.params_dtype = params_dtype
if quant_config is None:
self.quant_method: Optional[
QuantizeMethodBase] = UnquantizedLinearMethod()
self.quant_method: QuantizeMethodBase | None = UnquantizedLinearMethod(
)
else:
self.quant_method = quant_config.get_quant_method(self,
prefix=prefix)
def forward(self,
x: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
raise NotImplementedError
@@ -182,8 +180,8 @@ class ReplicatedLinear(LinearBase):
output_size: int,
bias: bool = True,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
super().__init__(input_size,
output_size,
@@ -223,8 +221,7 @@ class ReplicatedLinear(LinearBase):
f"to a parameter of size {param.size()}")
param.data.copy_(loaded_weight)
def forward(self,
x: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
bias = self.bias if not self.skip_bias_add else None
assert self.quant_method is not None
output = self.quant_method.apply(self, x, bias)
@@ -268,9 +265,9 @@ class ColumnParallelLinear(LinearBase):
bias: bool = True,
gather_output: bool = False,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
output_sizes: Optional[list[int]] = None,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
output_sizes: list[int] | None = None,
prefix: str = ""):
# Divide the weight matrix along the last dimension.
self.tp_size = get_tensor_model_parallel_world_size()
@@ -345,9 +342,8 @@ class ColumnParallelLinear(LinearBase):
loaded_weight = loaded_weight.reshape(1)
param.load_column_parallel_weight(loaded_weight=loaded_weight)
def forward(
self,
input_: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
def forward(self,
input_: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
bias = self.bias if not self.skip_bias_add else None
# Matrix multiply.
@@ -399,8 +395,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
bias: bool = True,
gather_output: bool = False,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
self.output_sizes = output_sizes
tp_size = get_tensor_model_parallel_world_size()
@@ -417,7 +413,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
def weight_loader(self,
param: Parameter,
loaded_weight: torch.Tensor,
loaded_shard_id: Optional[int] = None) -> None:
loaded_shard_id: int | None = None) -> None:
param_data = param.data
output_dim = getattr(param, "output_dim", None)
@@ -510,10 +506,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
# Special case for Quantization.
# If quantized, we need to adjust the offset and size to account
# for the packing.
if isinstance(
param,
(PackedColumnParameter,
PackedvLLMParameter)) and param.packed_dim == param.output_dim:
if isinstance(param, PackedColumnParameter | PackedvLLMParameter
) and param.packed_dim == param.output_dim:
shard_size, shard_offset = \
param.adjust_shard_indexes_for_packing(
shard_size=shard_size, shard_offset=shard_offset)
@@ -525,7 +519,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
def weight_loader_v2(self,
param: BasevLLMParameter,
loaded_weight: torch.Tensor,
loaded_shard_id: Optional[int] = None) -> None:
loaded_shard_id: int | None = None) -> None:
if loaded_shard_id is None:
if isinstance(param, PerTensorScaleParameter):
param.load_merged_column_weight(loaded_weight=loaded_weight,
@@ -598,11 +592,11 @@ class QKVParallelLinear(ColumnParallelLinear):
hidden_size: int,
head_size: int,
total_num_heads: int,
total_num_kv_heads: Optional[int] = None,
total_num_kv_heads: int | None = None,
bias: bool = True,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
self.hidden_size = hidden_size
self.head_size = head_size
@@ -637,7 +631,7 @@ class QKVParallelLinear(ColumnParallelLinear):
quant_config=quant_config,
prefix=prefix)
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> Optional[int]:
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> int | None:
shard_offset_mapping = {
"q": 0,
"k": self.num_heads * self.head_size,
@@ -646,7 +640,7 @@ class QKVParallelLinear(ColumnParallelLinear):
}
return shard_offset_mapping.get(loaded_shard_id)
def _get_shard_size_mapping(self, loaded_shard_id: str) -> Optional[int]:
def _get_shard_size_mapping(self, loaded_shard_id: str) -> int | None:
shard_size_mapping = {
"q": self.num_heads * self.head_size,
"k": self.num_kv_heads * self.head_size,
@@ -679,10 +673,8 @@ class QKVParallelLinear(ColumnParallelLinear):
# Special case for Quantization.
# If quantized, we need to adjust the offset and size to account
# for the packing.
if isinstance(
param,
(PackedColumnParameter,
PackedvLLMParameter)) and param.packed_dim == param.output_dim:
if isinstance(param, PackedColumnParameter | PackedvLLMParameter
) and param.packed_dim == param.output_dim:
shard_size, shard_offset = \
param.adjust_shard_indexes_for_packing(
shard_size=shard_size, shard_offset=shard_offset)
@@ -694,7 +686,7 @@ class QKVParallelLinear(ColumnParallelLinear):
def weight_loader_v2(self,
param: BasevLLMParameter,
loaded_weight: torch.Tensor,
loaded_shard_id: Optional[str] = None):
loaded_shard_id: str | None = None):
if loaded_shard_id is None: # special case for certain models
if isinstance(param, PerTensorScaleParameter):
param.load_qkv_weight(loaded_weight=loaded_weight, shard_id=0)
@@ -720,7 +712,7 @@ class QKVParallelLinear(ColumnParallelLinear):
def weight_loader(self,
param: Parameter,
loaded_weight: torch.Tensor,
loaded_shard_id: Optional[str] = None):
loaded_shard_id: str | None = None):
param_data = param.data
output_dim = getattr(param, "output_dim", None)
@@ -845,9 +837,9 @@ class RowParallelLinear(LinearBase):
bias: bool = True,
input_is_parallel: bool = True,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
params_dtype: torch.dtype | None = None,
reduce_results: bool = True,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
# Divide the weight matrix along the first dimension.
self.tp_rank = get_tensor_model_parallel_rank()
@@ -921,7 +913,7 @@ class RowParallelLinear(LinearBase):
param.load_row_parallel_weight(loaded_weight=loaded_weight)
def forward(self, input_) -> tuple[torch.Tensor, Optional[Parameter]]:
def forward(self, input_) -> tuple[torch.Tensor, Parameter | None]:
if self.input_is_parallel:
input_parallel = input_
else:
+2 -4
View File
@@ -1,7 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional
import torch
import torch.nn as nn
@@ -18,10 +16,10 @@ class MLP(nn.Module):
self,
input_dim: int,
mlp_hidden_dim: int,
output_dim: Optional[int] = None,
output_dim: int | None = None,
bias: bool = True,
act_type: str = "gelu_pytorch_tanh",
dtype: Optional[torch.dtype] = None,
dtype: torch.dtype | None = None,
prefix: str = "",
):
super().__init__()
@@ -3,7 +3,7 @@
import inspect
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Optional
from typing import TYPE_CHECKING, Any
import torch
from torch import nn
@@ -105,8 +105,8 @@ class QuantizationConfig(ABC):
raise NotImplementedError
@classmethod
def override_quantization_method(
cls, hf_quant_cfg, user_quant) -> Optional[QuantizationMethods]:
def override_quantization_method(cls, hf_quant_cfg,
user_quant) -> QuantizationMethods | None:
"""
Detects if this quantization method can support a given checkpoint
format by overriding the user specified quantization method --
@@ -135,7 +135,7 @@ class QuantizationConfig(ABC):
@abstractmethod
def get_quant_method(self, layer: torch.nn.Module,
prefix: str) -> Optional[QuantizeMethodBase]:
prefix: str) -> QuantizeMethodBase | None:
"""Get the quantize method to use for the quantized layer.
Args:
@@ -147,5 +147,5 @@ class QuantizationConfig(ABC):
"""
raise NotImplementedError
def get_cache_scale(self, name: str) -> Optional[str]:
return None
def get_cache_scale(self, name: str) -> str | None:
return None
+20 -20
View File
@@ -23,7 +23,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
"""Rotary Positional Embeddings."""
from typing import Any, Dict, List, Optional, Tuple, Union
from typing import Any
import torch
@@ -84,7 +84,7 @@ class RotaryEmbedding(CustomOp):
head_size: int,
rotary_dim: int,
max_position_embeddings: int,
base: Union[int, float],
base: int | float,
is_neox_style: bool,
dtype: torch.dtype,
) -> None:
@@ -101,7 +101,7 @@ class RotaryEmbedding(CustomOp):
self.cos_sin_cache: torch.Tensor
self.register_buffer("cos_sin_cache", cache, persistent=False)
def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor:
def _compute_inv_freq(self, base: int | float) -> torch.Tensor:
"""Compute the inverse frequency."""
# NOTE(woosuk): To exactly match the HF implementation, we need to
# use CPU to compute the cache and then move it to GPU. However, we
@@ -127,8 +127,8 @@ class RotaryEmbedding(CustomOp):
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
offsets: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
offsets: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""A PyTorch-native implementation of forward()."""
if offsets is not None:
positions = positions + offsets
@@ -159,7 +159,7 @@ class RotaryEmbedding(CustomOp):
return s
def _to_tuple(x: Union[int, Tuple[int, ...]], dim: int = 2) -> Tuple[int, ...]:
def _to_tuple(x: int | tuple[int, ...], dim: int = 2) -> tuple[int, ...]:
if isinstance(x, int):
return (x, ) * dim
elif len(x) == dim:
@@ -168,8 +168,8 @@ def _to_tuple(x: Union[int, Tuple[int, ...]], dim: int = 2) -> Tuple[int, ...]:
raise ValueError(f"Expected length {dim} or int, but got {x}")
def get_meshgrid_nd(start: Union[int, Tuple[int, ...]],
*args: Union[int, Tuple[int, ...]],
def get_meshgrid_nd(start: int | tuple[int, ...],
*args: int | tuple[int, ...],
dim: int = 2) -> torch.Tensor:
"""
Get n-D meshgrid with start, stop and num.
@@ -217,12 +217,12 @@ def get_meshgrid_nd(start: Union[int, Tuple[int, ...]],
def get_1d_rotary_pos_embed(
dim: int,
pos: Union[torch.FloatTensor, int],
pos: torch.FloatTensor | int,
theta: float = 10000.0,
theta_rescale_factor: float = 1.0,
interpolation_factor: float = 1.0,
dtype: torch.dtype = torch.float32,
) -> Tuple[torch.Tensor, torch.Tensor]:
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
(Note: `cis` means `cos + i * sin`, where i is the imaginary unit.)
@@ -261,13 +261,13 @@ def get_nd_rotary_pos_embed(
start,
*args,
theta=10000.0,
theta_rescale_factor: Union[float, List[float]] = 1.0,
interpolation_factor: Union[float, List[float]] = 1.0,
theta_rescale_factor: float | list[float] = 1.0,
interpolation_factor: float | list[float] = 1.0,
shard_dim: int = 0,
sp_rank: int = 0,
sp_world_size: int = 1,
dtype: torch.dtype = torch.float32,
) -> Tuple[torch.Tensor, torch.Tensor]:
) -> tuple[torch.Tensor, torch.Tensor]:
"""
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
Supports sequence parallelism by allowing sharding of a specific dimension.
@@ -324,7 +324,7 @@ def get_nd_rotary_pos_embed(
else:
grid = full_grid
if isinstance(theta_rescale_factor, (int, float)):
if isinstance(theta_rescale_factor, int | float):
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
elif isinstance(theta_rescale_factor,
list) and len(theta_rescale_factor) == 1:
@@ -333,7 +333,7 @@ def get_nd_rotary_pos_embed(
rope_dim_list
), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
if isinstance(interpolation_factor, (int, float)):
if isinstance(interpolation_factor, int | float):
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
elif isinstance(interpolation_factor,
list) and len(interpolation_factor) == 1:
@@ -370,7 +370,7 @@ def get_rotary_pos_embed(
interpolation_factor=1.0,
shard_dim: int = 0,
dtype: torch.dtype = torch.float32,
) -> Tuple[torch.Tensor, torch.Tensor]:
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Generate rotary positional embeddings for the given sizes.
@@ -417,17 +417,17 @@ def get_rotary_pos_embed(
return freqs_cos, freqs_sin
_ROPE_DICT: Dict[Tuple, RotaryEmbedding] = {}
_ROPE_DICT: dict[tuple, RotaryEmbedding] = {}
def get_rope(
head_size: int,
rotary_dim: int,
max_position: int,
base: Union[int, float],
base: int | float,
is_neox_style: bool = True,
rope_scaling: Optional[Dict[str, Any]] = None,
dtype: Optional[torch.dtype] = None,
rope_scaling: dict[str, Any] | None = None,
dtype: torch.dtype | None = None,
partial_rotary_factor: float = 1.0,
) -> RotaryEmbedding:
if dtype is None:
+1 -2
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/utils.py
"""Utility methods for model layers."""
from typing import Tuple
import torch
@@ -10,7 +9,7 @@ def get_token_bin_counts_and_mask(
tokens: torch.Tensor,
vocab_size: int,
num_seqs: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
) -> tuple[torch.Tensor, torch.Tensor]:
# Compute the bin counts for the tokens.
# vocab_size + 1 for padding.
bin_counts = torch.zeros((num_seqs, vocab_size + 1),
+2 -3
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Optional
import torch
import torch.nn as nn
@@ -36,7 +35,7 @@ class PatchEmbed(nn.Module):
prefix: str = ""):
super().__init__()
# Convert patch_size to 2-tuple
if isinstance(patch_size, (list, tuple)):
if isinstance(patch_size, list | tuple):
if len(patch_size) == 1:
patch_size = (patch_size[0], patch_size[0])
else:
@@ -133,7 +132,7 @@ class ModulateProjection(nn.Module):
hidden_size: int,
factor: int = 2,
act_layer: str = "silu",
dtype: Optional[torch.dtype] = None,
dtype: torch.dtype | None = None,
prefix: str = "",
):
super().__init__()
+11 -11
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Sequence
from dataclasses import dataclass
from typing import List, Optional, Sequence, Tuple
import torch
import torch.nn.functional as F
@@ -24,7 +24,7 @@ class UnquantizedEmbeddingMethod(QuantizeMethodBase):
def create_weights(self, layer: torch.nn.Module,
input_size_per_partition: int,
output_partition_sizes: List[int], input_size: int,
output_partition_sizes: list[int], input_size: int,
output_size: int, params_dtype: torch.dtype,
**extra_weight_attrs):
"""Create weights for embedding layer."""
@@ -39,7 +39,7 @@ class UnquantizedEmbeddingMethod(QuantizeMethodBase):
def apply(self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None) -> torch.Tensor:
bias: torch.Tensor | None = None) -> torch.Tensor:
return F.linear(x, layer.weight, bias)
def embedding(self, layer: torch.nn.Module,
@@ -139,7 +139,7 @@ def get_masked_input_and_mask(
input_: torch.Tensor, org_vocab_start_index: int,
org_vocab_end_index: int, num_org_vocab_padding: int,
added_vocab_start_index: int,
added_vocab_end_index: int) -> Tuple[torch.Tensor, torch.Tensor]:
added_vocab_end_index: int) -> tuple[torch.Tensor, torch.Tensor]:
# torch.compile will fuse all of the pointwise ops below
# into a single kernel, making it very fast
org_vocab_mask = (input_ >= org_vocab_start_index) & (input_
@@ -197,10 +197,10 @@ class VocabParallelEmbedding(torch.nn.Module):
def __init__(self,
num_embeddings: int,
embedding_dim: int,
params_dtype: Optional[torch.dtype] = None,
org_num_embeddings: Optional[int] = None,
params_dtype: torch.dtype | None = None,
org_num_embeddings: int | None = None,
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
super().__init__()
@@ -296,7 +296,7 @@ class VocabParallelEmbedding(torch.nn.Module):
org_vocab_start_index, org_vocab_end_index, added_vocab_start_index,
added_vocab_end_index)
def get_sharded_to_full_mapping(self) -> Optional[List[int]]:
def get_sharded_to_full_mapping(self) -> list[int] | None:
"""Get a mapping that can be used to reindex the gathered
logits for sampling.
@@ -310,9 +310,9 @@ class VocabParallelEmbedding(torch.nn.Module):
if self.tp_size < 2:
return None
base_embeddings: List[int] = []
added_embeddings: List[int] = []
padding: List[int] = []
base_embeddings: list[int] = []
added_embeddings: list[int] = []
padding: list[int] = []
for tp_rank in range(self.tp_size):
shard_indices = self._get_indices(self.num_embeddings_padded,
self.org_vocab_size_padded,
+2 -3
View File
@@ -11,7 +11,7 @@ from logging import Logger
from logging.config import dictConfig
from os import path
from types import MethodType
from typing import Any, Optional, cast
from typing import Any, cast
import fastvideo.v1.envs as envs
@@ -278,8 +278,7 @@ def _trace_calls(log_path, root_dir, frame, event, arg=None):
return partial(_trace_calls, log_path, root_dir)
def enable_trace_function_call(log_file_path: str,
root_dir: Optional[str] = None):
def enable_trace_function_call(log_file_path: str, root_dir: str | None = None):
"""
Enable tracing of every function call in code under `root_dir`.
This is useful for debugging hangs or crashes.
+8 -10
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import Any, List, Optional, Tuple, Union
from typing import Any
import torch
from torch import nn
@@ -18,7 +18,7 @@ class BaseDiT(nn.Module, ABC):
num_attention_heads: int
num_channels_latents: int
# always supports torch_sdpa
_supported_attention_backends: Tuple[
_supported_attention_backends: tuple[
_Backend, ...] = DiTConfig()._supported_attention_backends
def __init_subclass__(cls) -> None:
@@ -33,11 +33,9 @@ class BaseDiT(nn.Module, ABC):
f"Subclasses of BaseDiT must define '{attr}' class variable"
)
def __init__(self, config: DiTConfig, hf_config: dict[str, Any],
**kwargs) -> None:
def __init__(self, config: DiTConfig, **kwargs) -> None:
super().__init__()
self.config = config
self.hf_config = hf_config
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
@@ -46,10 +44,10 @@ class BaseDiT(nn.Module, ABC):
@abstractmethod
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance=None,
**kwargs) -> torch.Tensor:
pass
@@ -65,7 +63,7 @@ class BaseDiT(nn.Module, ABC):
)
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
def supported_attention_backends(self) -> tuple[_Backend, ...]:
return self._supported_attention_backends
@@ -83,7 +81,7 @@ class CachableDiT(BaseDiT):
num_attention_heads: int
num_channels_latents: int
# always supports torch_sdpa
_supported_attention_backends: Tuple[
_supported_attention_backends: tuple[
_Backend, ...] = DiTConfig()._supported_attention_backends
def __init__(self, config: DiTConfig, **kwargs) -> None:
+11 -13
View File
@@ -1,7 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Any, List, Optional, Tuple, Union
import numpy as np
import torch
import torch.nn as nn
@@ -96,8 +94,8 @@ class MMDoubleStreamBlock(nn.Module):
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[_Backend, ...] | None = None,
prefix: str = "",
):
super().__init__()
@@ -202,7 +200,7 @@ class MMDoubleStreamBlock(nn.Module):
txt: torch.Tensor,
vec: torch.Tensor,
freqs_cis: tuple,
) -> Tuple[torch.Tensor, torch.Tensor]:
) -> tuple[torch.Tensor, torch.Tensor]:
# Process modulation vectors
img_mod_outputs = self.img_mod(vec)
(
@@ -303,8 +301,8 @@ class MMSingleStreamBlock(nn.Module):
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[_Backend, ...] | None = None,
prefix: str = "",
):
super().__init__()
@@ -366,7 +364,7 @@ class MMSingleStreamBlock(nn.Module):
x: torch.Tensor,
vec: torch.Tensor,
txt_len: int,
freqs_cis: Tuple[torch.Tensor, torch.Tensor],
freqs_cis: tuple[torch.Tensor, torch.Tensor],
) -> torch.Tensor:
# Process modulation
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
@@ -442,8 +440,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
)._supported_attention_backends
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
def __init__(self, config: HunyuanVideoConfig, hf_config: dict[str, Any]):
super().__init__(config=config, hf_config=hf_config)
def __init__(self, config: HunyuanVideoConfig):
super().__init__(config=config)
self.patch_size = [
config.patch_size_t, config.patch_size, config.patch_size
@@ -542,10 +540,10 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
# TODO: change output to a dict
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance=None,
**kwargs):
"""
+19 -19
View File
@@ -10,7 +10,6 @@
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
from typing import Any, Dict, Optional, Tuple
import torch
from einops import rearrange, repeat
@@ -55,7 +54,7 @@ class PatchEmbed2D(nn.Module):
prefix: str = ""):
super().__init__()
# Convert patch_size to 2-tuple
if isinstance(patch_size, (list, tuple)):
if isinstance(patch_size, list | tuple):
if len(patch_size) == 1:
patch_size = (patch_size[0], patch_size[0])
else:
@@ -143,7 +142,7 @@ class SelfAttention(nn.Module):
def __init__(self,
hidden_dim,
head_dim,
rope_split: Tuple[int, int, int] = (64, 32, 32),
rope_split: tuple[int, int, int] = (64, 32, 32),
bias: bool = False,
with_rope: bool = True,
with_qk_norm: bool = True,
@@ -190,8 +189,10 @@ class SelfAttention(nn.Module):
outs = []
idx = 0
for (chunk_size, cos_i, sin_i) in zip(self.rope_split, cos_splits,
sin_splits):
for (chunk_size, cos_i, sin_i) in zip(self.rope_split,
cos_splits,
sin_splits,
strict=False):
# slice the corresponding channels
x_chunk = x[..., idx:idx + chunk_size] # [B,S,H,chunk_size]
idx += chunk_size
@@ -331,8 +332,8 @@ class AdaLayerNormSingle(nn.Module):
def forward(
self,
timestep: torch.Tensor,
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
added_cond_kwargs: dict[str, torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
embedded_timestep = self.emb(timestep * self.time_step_rescale)
out, _ = self.linear(self.silu(embedded_timestep))
@@ -377,7 +378,7 @@ class StepVideoTransformerBlock(nn.Module):
dim: int,
attention_head_dim: int,
norm_eps: float = 1e-5,
ff_inner_dim: Optional[int] = None,
ff_inner_dim: int | None = None,
ff_bias: bool = False,
attention_type: str = 'torch'):
super().__init__()
@@ -417,7 +418,7 @@ class StepVideoTransformerBlock(nn.Module):
kv: torch.Tensor,
t_expand: torch.LongTensor,
attn_mask=None,
rope_positions: Optional[list] = None,
rope_positions: list | None = None,
cos_sin=None,
mask_strategy=None) -> torch.Tensor:
@@ -462,9 +463,8 @@ class StepVideoModel(BaseDiT):
_supported_attention_backends = StepVideoConfig(
)._supported_attention_backends
def __init__(self, config: StepVideoConfig, hf_config: dict[str,
Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
def __init__(self, config: StepVideoConfig) -> None:
super().__init__(config=config)
self.num_attention_heads = config.num_attention_heads
self.attention_head_dim = config.attention_head_dim
self.in_channels = config.in_channels
@@ -539,7 +539,7 @@ class StepVideoModel(BaseDiT):
return hidden_states
def prepare_attn_mask(self, encoder_attention_mask, encoder_hidden_states,
q_seqlen) -> Tuple[torch.Tensor, torch.Tensor]:
q_seqlen) -> tuple[torch.Tensor, torch.Tensor]:
kv_seqlens = encoder_attention_mask.sum(dim=1).int()
mask = torch.zeros([len(kv_seqlens), q_seqlen,
max(kv_seqlens)],
@@ -594,12 +594,12 @@ class StepVideoModel(BaseDiT):
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
t_expand: Optional[torch.LongTensor] = None,
encoder_hidden_states_2: Optional[torch.Tensor] = None,
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
encoder_attention_mask: Optional[torch.Tensor] = None,
fps: Optional[torch.Tensor] = None,
encoder_hidden_states: torch.Tensor | None = None,
t_expand: torch.LongTensor | None = None,
encoder_hidden_states_2: torch.Tensor | None = None,
added_cond_kwargs: dict[str, torch.Tensor] | None = None,
encoder_attention_mask: torch.Tensor | None = None,
fps: torch.Tensor | None = None,
return_dict: bool = True,
mask_strategy=None,
guidance=None,
+13 -15
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any, List, Optional, Tuple, Union
import numpy as np
import torch
@@ -53,7 +52,7 @@ class WanTimeTextImageEmbedding(nn.Module):
dim: int,
time_freq_dim: int,
text_embed_dim: int,
image_embed_dim: Optional[int] = None,
image_embed_dim: int | None = None,
):
super().__init__()
@@ -76,7 +75,7 @@ class WanTimeTextImageEmbedding(nn.Module):
self,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
encoder_hidden_states_image: torch.Tensor | None = None,
):
temb = self.time_embedder(timestep)
timestep_proj = self.time_modulation(temb)
@@ -173,7 +172,7 @@ class WanI2VCrossAttention(WanSelfAttention):
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
supported_attention_backends: tuple[_Backend, ...] | None = None
) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps,
supported_attention_backends)
@@ -222,9 +221,9 @@ class WanTransformerBlock(nn.Module):
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[_Backend, ...]
| None = None,
prefix: str = ""):
super().__init__()
@@ -292,13 +291,13 @@ class WanTransformerBlock(nn.Module):
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: Tuple[torch.Tensor, torch.Tensor],
freqs_cis: tuple[torch.Tensor, torch.Tensor],
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
@@ -360,9 +359,8 @@ class WanTransformer3DModel(CachableDiT):
)._supported_attention_backends
_param_names_mapping = WanVideoConfig()._param_names_mapping
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
def __init__(self, config: WanVideoConfig) -> None:
super().__init__(config=config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
@@ -418,10 +416,10 @@ class WanTransformer3DModel(CachableDiT):
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance=None,
**kwargs) -> torch.Tensor:
forward_batch = get_forward_context().forward_batch
+9 -10
View File
@@ -1,5 +1,4 @@
from abc import ABC, abstractmethod
from typing import Optional, Tuple
import torch
from torch import nn
@@ -11,7 +10,7 @@ from fastvideo.v1.platforms import _Backend
class TextEncoder(nn.Module, ABC):
_supported_attention_backends: Tuple[
_supported_attention_backends: tuple[
_Backend, ...] = TextEncoderConfig()._supported_attention_backends
def __init__(self, config: TextEncoderConfig) -> None:
@@ -24,21 +23,21 @@ class TextEncoder(nn.Module, ABC):
@abstractmethod
def forward(self,
input_ids: Optional[torch.Tensor],
position_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.Tensor] = None,
output_hidden_states: Optional[bool] = None,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs) -> BaseEncoderOutput:
pass
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
def supported_attention_backends(self) -> tuple[_Backend, ...]:
return self._supported_attention_backends
class ImageEncoder(nn.Module, ABC):
_supported_attention_backends: Tuple[
_supported_attention_backends: tuple[
_Backend, ...] = ImageEncoderConfig()._supported_attention_backends
def __init__(self, config: ImageEncoderConfig) -> None:
@@ -55,5 +54,5 @@ class ImageEncoder(nn.Module, ABC):
pass
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
def supported_attention_backends(self) -> tuple[_Backend, ...]:
return self._supported_attention_backends
+38 -38
View File
@@ -3,7 +3,7 @@
# Adapted from transformers: https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py
"""Minimal implementation of CLIPVisionModel intended to be only used
within a vision language model."""
from typing import Iterable, Optional, Set, Tuple, Union
from collections.abc import Iterable
import torch
import torch.nn as nn
@@ -91,9 +91,9 @@ class CLIPTextEmbeddings(nn.Module):
def forward(
self,
input_ids: Optional[torch.LongTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
input_ids: torch.LongTensor | None = None,
position_ids: torch.LongTensor | None = None,
inputs_embeds: torch.FloatTensor | None = None,
) -> torch.Tensor:
if input_ids is not None:
seq_length = input_ids.shape[-1]
@@ -128,8 +128,8 @@ class CLIPAttention(nn.Module):
def __init__(
self,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
config: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__()
@@ -209,8 +209,8 @@ class CLIPMLP(nn.Module):
def __init__(
self,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
config: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -239,8 +239,8 @@ class CLIPEncoderLayer(nn.Module):
def __init__(
self,
config: Union[CLIPTextConfig, CLIPVisionConfig],
quant_config: Optional[QuantizationConfig] = None,
config: CLIPTextConfig | CLIPVisionConfig,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -284,9 +284,9 @@ class CLIPEncoder(nn.Module):
def __init__(
self,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = None,
config: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -305,8 +305,8 @@ class CLIPEncoder(nn.Module):
])
def forward(
self, inputs_embeds: torch.Tensor, return_all_hidden_states: bool
) -> Union[torch.Tensor, list[torch.Tensor]]:
self, inputs_embeds: torch.Tensor, return_all_hidden_states: bool
) -> torch.Tensor | list[torch.Tensor]:
hidden_states_pool = [inputs_embeds]
hidden_states = inputs_embeds
@@ -325,8 +325,8 @@ class CLIPTextTransformer(nn.Module):
def __init__(self,
config: CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = None,
quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = None,
prefix: str = ""):
super().__init__()
self.config = config
@@ -348,11 +348,11 @@ class CLIPTextTransformer(nn.Module):
def forward(
self,
input_ids: Optional[torch.Tensor],
position_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.Tensor] = None,
output_hidden_states: Optional[bool] = None,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
) -> BaseEncoderOutput:
r"""
Returns:
@@ -440,11 +440,11 @@ class CLIPTextModel(TextEncoder):
def forward(
self,
input_ids: Optional[torch.Tensor],
position_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.Tensor] = None,
output_hidden_states: Optional[bool] = None,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
@@ -456,8 +456,8 @@ class CLIPTextModel(TextEncoder):
)
return outputs
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
def load_weights(self, weights: Iterable[tuple[str,
torch.Tensor]]) -> set[str]:
# Define mapping for stacked parameters
stacked_params_mapping = [
@@ -467,7 +467,7 @@ class CLIPTextModel(TextEncoder):
("qkv_proj", "v_proj", "v"),
]
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
loaded_params: set[str] = set()
for name, loaded_weight in weights:
# Handle q_proj, k_proj, v_proj -> qkv_proj mapping
for param_name, weight_name, shard_id in stacked_params_mapping:
@@ -498,9 +498,9 @@ class CLIPVisionTransformer(nn.Module):
def __init__(
self,
config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = None,
require_post_norm: Optional[bool] = None,
quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = None,
require_post_norm: bool | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -540,7 +540,7 @@ class CLIPVisionTransformer(nn.Module):
def forward(
self,
pixel_values: torch.Tensor,
feature_sample_layers: Optional[list[int]] = None,
feature_sample_layers: list[int] | None = None,
) -> torch.Tensor:
hidden_states = self.embeddings(pixel_values)
@@ -582,7 +582,7 @@ class CLIPVisionModel(ImageEncoder):
def forward(
self,
pixel_values: torch.Tensor,
feature_sample_layers: Optional[list[int]] = None,
feature_sample_layers: list[int] | None = None,
**kwargs,
) -> BaseEncoderOutput:
last_hidden_state = self.vision_model(pixel_values,
@@ -595,8 +595,8 @@ class CLIPVisionModel(ImageEncoder):
# (TODO) Add prefix argument for filtering out weights to be loaded
# ref: https://github.com/vllm-project/vllm/pull/7186#discussion_r1734163986
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
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"),
@@ -604,7 +604,7 @@ class CLIPVisionModel(ImageEncoder):
("qkv_proj", "v_proj", "v"),
]
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
loaded_params: set[str] = set()
layer_count = len(self.vision_model.encoder.layers)
for name, loaded_weight in weights:
+18 -17
View File
@@ -23,7 +23,8 @@
# See the License for the specific language governing permissions and
# limitations under the License.
"""Inference-only LLaMA model compatible with HuggingFace weights."""
from typing import Any, Dict, Iterable, Optional, Set, Tuple
from collections.abc import Iterable
from typing import Any
import torch
from torch import nn
@@ -52,7 +53,7 @@ class LlamaMLP(nn.Module):
hidden_size: int,
intermediate_size: int,
hidden_act: str,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
bias: bool = False,
prefix: str = "",
) -> None:
@@ -92,9 +93,9 @@ class LlamaAttention(nn.Module):
num_heads: int,
num_kv_heads: int,
rope_theta: float = 10000,
rope_scaling: Optional[Dict[str, Any]] = None,
rope_scaling: dict[str, Any] | None = None,
max_position_embeddings: int = 8192,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
bias: bool = False,
bias_o_proj: bool = False,
prefix: str = "") -> None:
@@ -201,7 +202,7 @@ class LlamaDecoderLayer(nn.Module):
def __init__(
self,
config: LlamaConfig,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -254,8 +255,8 @@ class LlamaDecoderLayer(nn.Module):
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
residual: Optional[torch.Tensor],
) -> Tuple[torch.Tensor, torch.Tensor]:
residual: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
# Self Attention
if residual is None:
residual = hidden_states
@@ -318,11 +319,11 @@ class LlamaModel(TextEncoder):
def forward(
self,
input_ids: Optional[torch.Tensor],
position_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.Tensor] = None,
output_hidden_states: Optional[bool] = None,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
output_hidden_states = (output_hidden_states
@@ -339,7 +340,7 @@ class LlamaModel(TextEncoder):
0, hidden_states.shape[1],
device=hidden_states.device).unsqueeze(0)
all_hidden_states: Optional[Tuple[Any, ...]] = (
all_hidden_states: tuple[Any, ...] | None = (
) if output_hidden_states else None
for layer in self.layers:
if all_hidden_states is not None:
@@ -367,8 +368,8 @@ class LlamaModel(TextEncoder):
return output
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
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"),
@@ -378,7 +379,7 @@ class LlamaModel(TextEncoder):
(".gate_up_proj", ".up_proj", 1),
]
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
loaded_params: set[str] = set()
for name, loaded_weight in weights:
if "rotary_emb.inv_freq" in name:
continue
@@ -400,7 +401,7 @@ class LlamaModel(TextEncoder):
# continue
if "scale" in name:
# Remapping the name of FP8 kv-scale.
kv_scale_name: Optional[str] = maybe_remap_kv_scale_name(
kv_scale_name: str | None = maybe_remap_kv_scale_name(
name, params_dict)
if kv_scale_name is None:
continue
+8 -9
View File
@@ -13,7 +13,6 @@
# ==============================================================================
import os
from functools import wraps
from typing import List, Optional
import torch
import torch.nn as nn
@@ -179,10 +178,10 @@ class StepChatTokenizer:
def vocab_size(self):
return self._tokenizer.vocab_size()
def tokenize(self, text: str) -> List[int]:
def tokenize(self, text: str) -> list[int]:
return self._tokenizer.encode_as_ids(text)
def detokenize(self, token_ids: List[int]) -> str:
def detokenize(self, token_ids: list[int]) -> str:
return self._tokenizer.decode_ids(token_ids)
@@ -347,9 +346,9 @@ class MultiQueryAttention(nn.Module):
def forward(
self,
x: torch.Tensor,
mask: Optional[torch.Tensor],
cu_seqlens: Optional[torch.Tensor],
max_seq_len: Optional[torch.Tensor],
mask: torch.Tensor | None,
cu_seqlens: torch.Tensor | None,
max_seq_len: torch.Tensor | None,
):
seqlen, bsz, dim = x.shape
xqkv = self.wqkv(x)
@@ -471,9 +470,9 @@ class TransformerBlock(nn.Module):
def forward(
self,
x: torch.Tensor,
mask: Optional[torch.Tensor],
cu_seqlens: Optional[torch.Tensor],
max_seq_len: Optional[torch.Tensor],
mask: torch.Tensor | None,
cu_seqlens: torch.Tensor | None,
max_seq_len: torch.Tensor | None,
):
residual = self.attention.forward(self.attention_norm(x), mask,
cu_seqlens, max_seq_len)
+29 -29
View File
@@ -20,8 +20,8 @@
"""PyTorch T5 & UMT5 model."""
import math
from collections.abc import Iterable
from dataclasses import dataclass
from typing import Iterable, Optional, Set, Tuple
import torch
import torch.nn.functional as F
@@ -64,7 +64,7 @@ class T5DenseActDense(nn.Module):
def __init__(self,
config: T5Config,
quant_config: Optional[QuantizationConfig] = None):
quant_config: QuantizationConfig | None = None):
super().__init__()
self.wi = MergedColumnParallelLinear(config.d_model, [config.d_ff],
bias=False)
@@ -85,7 +85,7 @@ class T5DenseGatedActDense(nn.Module):
def __init__(self,
config: T5Config,
quant_config: Optional[QuantizationConfig] = None):
quant_config: QuantizationConfig | None = None):
super().__init__()
self.wi_0 = MergedColumnParallelLinear(config.d_model, [config.d_ff],
bias=False,
@@ -113,7 +113,7 @@ class T5LayerFF(nn.Module):
def __init__(self,
config: T5Config,
quant_config: Optional[QuantizationConfig] = None):
quant_config: QuantizationConfig | None = None):
super().__init__()
if config.is_gated_act:
self.DenseReluDense = T5DenseGatedActDense(
@@ -155,7 +155,7 @@ class T5Attention(nn.Module):
config: T5Config,
attn_type: str,
has_relative_attention_bias=False,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
super().__init__()
self.attn_type = attn_type
@@ -294,7 +294,7 @@ class T5Attention(nn.Module):
self,
hidden_states: torch.Tensor, # (num_tokens, d_model)
attention_mask: torch.Tensor,
attn_metadata: Optional[AttentionMetadata] = None,
attn_metadata: AttentionMetadata | None = None,
) -> torch.Tensor:
bs, seq_len, _ = hidden_states.shape
num_seqs = bs
@@ -344,7 +344,7 @@ class T5LayerSelfAttention(nn.Module):
self,
config,
has_relative_attention_bias=False,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__()
@@ -361,7 +361,7 @@ class T5LayerSelfAttention(nn.Module):
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
attn_metadata: Optional[AttentionMetadata] = None,
attn_metadata: AttentionMetadata | None = None,
) -> torch.Tensor:
normed_hidden_states = self.layer_norm.forward_native(hidden_states)
attention_output = self.SelfAttention(
@@ -377,7 +377,7 @@ class T5LayerCrossAttention(nn.Module):
def __init__(self,
config,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
super().__init__()
self.EncDecAttention = T5Attention(config,
@@ -390,7 +390,7 @@ class T5LayerCrossAttention(nn.Module):
def forward(
self,
hidden_states: torch.Tensor,
attn_metadata: Optional[AttentionMetadata] = None,
attn_metadata: AttentionMetadata | None = None,
) -> torch.Tensor:
normed_hidden_states = self.layer_norm.forward_native(hidden_states)
attention_output = self.EncDecAttention(
@@ -407,7 +407,7 @@ class T5Block(nn.Module):
config: T5Config,
is_decoder: bool,
has_relative_attention_bias=False,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
super().__init__()
self.is_decoder = is_decoder
@@ -431,7 +431,7 @@ class T5Block(nn.Module):
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
attn_metadata: Optional[AttentionMetadata] = None,
attn_metadata: AttentionMetadata | None = None,
) -> torch.Tensor:
hidden_states = self.layer[0](hidden_states=hidden_states,
@@ -455,7 +455,7 @@ class T5Stack(nn.Module):
is_decoder: bool,
n_layers: int,
embed_tokens=None,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
is_umt5: bool = False):
super().__init__()
@@ -524,11 +524,11 @@ class T5EncoderModel(TextEncoder):
def forward(
self,
input_ids: Optional[torch.Tensor],
position_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.Tensor] = None,
output_hidden_states: Optional[bool] = None,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
attn_metadata = AttentionMetadata(None)
@@ -540,8 +540,8 @@ class T5EncoderModel(TextEncoder):
return BaseEncoderOutput(last_hidden_state=hidden_states)
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
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"),
@@ -549,7 +549,7 @@ class T5EncoderModel(TextEncoder):
(".qkv_proj", ".v", "v"),
]
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
loaded_params: set[str] = set()
for name, loaded_weight in weights:
loaded = False
if "decoder" in name or "lm_head" in name:
@@ -611,11 +611,11 @@ class UMT5EncoderModel(TextEncoder):
def forward(
self,
input_ids: Optional[torch.Tensor],
position_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.Tensor] = None,
output_hidden_states: Optional[bool] = None,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
attn_metadata = AttentionMetadata(None)
@@ -630,8 +630,8 @@ class UMT5EncoderModel(TextEncoder):
attention_mask=attention_mask,
)
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
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"),
@@ -639,7 +639,7 @@ class UMT5EncoderModel(TextEncoder):
(".qkv_proj", ".v", "v"),
]
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
loaded_params: set[str] = set()
for name, loaded_weight in weights:
loaded = False
if "decoder" in name or "lm_head" in name:
+4 -4
View File
@@ -2,7 +2,7 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/models/vision.py
from abc import ABC, abstractmethod
from typing import Generic, Optional, TypeVar, Union
from typing import Generic, TypeVar
import torch
from transformers import PretrainedConfig
@@ -48,9 +48,9 @@ class VisionEncoderInfo(ABC, Generic[_C]):
def resolve_visual_encoder_outputs(
encoder_outputs: Union[torch.Tensor, list[torch.Tensor]],
feature_sample_layers: Optional[list[int]],
post_layer_norm: Optional[torch.nn.LayerNorm],
encoder_outputs: torch.Tensor | list[torch.Tensor],
feature_sample_layers: list[int] | None,
post_layer_norm: torch.nn.LayerNorm | None,
max_possible_layers: int,
) -> torch.Tensor:
"""Given the outputs a visual encoder module that may correspond to the
+8 -8
View File
@@ -20,14 +20,14 @@ import contextlib
import json
import os
from pathlib import Path
from typing import Any, Dict, Optional, Type, Union
from typing import Any
from huggingface_hub import snapshot_download
from transformers import AutoConfig, PretrainedConfig
from transformers.models.auto.modeling_auto import (
MODEL_FOR_CAUSAL_LM_MAPPING_NAMES)
_CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
_CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = {
# ChatGLMConfig.model_type: ChatGLMConfig,
# DbrxConfig.model_type: DbrxConfig,
# ExaoneConfig.model_type: ExaoneConfig,
@@ -50,8 +50,8 @@ def download_from_hf(model_path: str):
def get_hf_config(
model: str,
trust_remote_code: bool,
revision: Optional[str] = None,
model_override_args: Optional[dict] = None,
revision: str | None = None,
model_override_args: dict | None = None,
**kwargs,
):
is_gguf = check_gguf_file(model)
@@ -83,8 +83,8 @@ def get_hf_config(
def get_diffusers_config(
model: str,
fastvideo_args: Optional[dict] = None,
) -> Dict[str, Any]:
fastvideo_args: dict | None = None,
) -> dict[str, Any]:
"""Gets a configuration for the given diffusers model.
Args:
@@ -104,7 +104,7 @@ def get_diffusers_config(
try:
# Load the config directly from the file
with open(config_file) as f:
config_dict: Dict[str, Any] = json.load(f)
config_dict: dict[str, Any] = json.load(f)
# TODO(will): apply any overrides from inference args
return config_dict
@@ -139,7 +139,7 @@ def attach_additional_stop_token_ids(tokenizer):
tokenizer.additional_stop_token_ids = None
def check_gguf_file(model: Union[str, os.PathLike]) -> bool:
def check_gguf_file(model: str | os.PathLike) -> bool:
"""Check if the file is a GGUF model."""
model = Path(model)
if not model.is_file():
+11 -29
View File
@@ -6,8 +6,8 @@ import json
import os
import time
from abc import ABC, abstractmethod
from copy import deepcopy
from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
from collections.abc import Generator, Iterable
from typing import Any, cast
import torch
import torch.nn as nn
@@ -106,7 +106,7 @@ class TextEncoderLoader(ComponentLoader):
fall_back_to_pt: bool = True
"""Whether .pt weights can be used."""
allow_patterns_overrides: Optional[list[str]] = None
allow_patterns_overrides: list[str] | None = None
"""If defined, weights will load exclusively using these patterns."""
counter_before_loading_weights: float = 0.0
@@ -116,8 +116,8 @@ class TextEncoderLoader(ComponentLoader):
self,
model_name_or_path: str,
fall_back_to_pt: bool,
allow_patterns_overrides: Optional[list[str]],
) -> Tuple[str, List[str], bool]:
allow_patterns_overrides: list[str] | None,
) -> tuple[str, list[str], bool]:
"""Prepare weights for the model.
If the model is not local, it will be downloaded."""
@@ -139,7 +139,7 @@ class TextEncoderLoader(ComponentLoader):
hf_folder = model_name_or_path
hf_weights_files: List[str] = []
hf_weights_files: list[str] = []
for pattern in allow_patterns:
hf_weights_files += glob.glob(os.path.join(hf_folder, pattern))
if len(hf_weights_files) > 0:
@@ -162,7 +162,7 @@ class TextEncoderLoader(ComponentLoader):
def _get_weights_iterator(
self, source: "Source"
) -> Generator[Tuple[str, torch.Tensor], None, None]:
) -> 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,
@@ -182,7 +182,7 @@ class TextEncoderLoader(ComponentLoader):
self,
model_config: Any,
model: nn.Module,
) -> Generator[Tuple[str, torch.Tensor], None, None]:
) -> Generator[tuple[str, torch.Tensor], None, None]:
primary_weights = TextEncoderLoader.Source(
model_config.model,
prefix="",
@@ -367,7 +367,6 @@ class TransformerLoader(ComponentLoader):
fastvideo_args: FastVideoArgs):
"""Load the transformer based on the model path, architecture, and inference args."""
config = get_diffusers_config(model=model_path)
hf_config = deepcopy(config)
cls_name = config.pop("_class_name")
if cls_name is None:
raise ValueError(
@@ -394,30 +393,13 @@ class TransformerLoader(ComponentLoader):
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
# Load the model using FSDP loader
logger.info("Loading model from %s, default_dtype: %s", cls_name, default_dtype)
# model = load_fsdp_model(model_cls=model_cls,
# init_params={
# "config": dit_config,
# "hf_config": hf_config
# },
# weight_dir_list=safetensors_list,
# device=fastvideo_args.device,
# cpu_offload=fastvideo_args.use_cpu_offload,
# default_dtype=default_dtype)
logger.info("Loading model from %s", cls_name)
model = load_fsdp_model(model_cls=model_cls,
init_params={
"config": dit_config,
"hf_config": hf_config
},
init_params={"config": dit_config},
weight_dir_list=safetensors_list,
device=fastvideo_args.device,
cpu_offload=fastvideo_args.use_cpu_offload,
default_dtype=default_dtype,
# TODO(will): make these configurable
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
)
default_dtype=default_dtype)
if fastvideo_args.enable_torch_compile:
logger.info("Torch Compile enabled for DiT")
for n, m in reversed(list(model.named_modules())):
+14 -39
View File
@@ -7,23 +7,20 @@
import contextlib
import re
from collections import defaultdict
from collections.abc import Callable, Generator, Hashable
from itertools import chain
from typing import (Any, Callable, DefaultDict, Dict, Generator, Hashable, List,
Optional, Tuple, Type)
from typing import Any
import torch
from torch import nn
from torch.distributed import DeviceMesh, init_device_mesh
from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard, MixedPrecisionPolicy
from torch.distributed._composable.fsdp import CPUOffloadPolicy, fully_shard
from torch.distributed._tensor import distribute_tensor
from torch.nn.modules.module import _IncompatibleKeys
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
# TODO(PY): move this to utils elsewhere
@@ -55,7 +52,7 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
def get_param_names_mapping(
mapping_dict: Dict[str, str]) -> Callable[[str], tuple[str, Any, Any]]:
mapping_dict: dict[str, str]) -> Callable[[str], tuple[str, Any, Any]]:
"""
Creates a mapping function that transforms parameter names using regex patterns.
@@ -89,29 +86,16 @@ def get_param_names_mapping(
# TODO(PY): add compile option
# param_dtype: torch.dtype,
# reduce_dtype: torch.dtype,
# output_dtype: torch.dtype,
# pp_enabled: bool = False,
# cpu_offload: bool = False,
def load_fsdp_model(
model_cls: Type[nn.Module],
init_params: Dict[str, Any],
weight_dir_list: List[str],
model_cls: type[nn.Module],
init_params: dict[str, Any],
weight_dir_list: list[str],
device: torch.device,
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
cpu_offload: bool = False,
output_dtype: Optional[torch.dtype] = None,
default_dtype: torch.dtype | None = torch.bfloat16,
) -> torch.nn.Module:
mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=True)
# with set_default_dtype(default_dtype), torch.device("meta"):
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
device_mesh = init_device_mesh(
"cuda",
mesh_shape=(get_sequence_model_parallel_world_size(), ),
@@ -120,7 +104,6 @@ def load_fsdp_model(
shard_model(model,
cpu_offload=cpu_offload,
reshard_after_forward=True,
mp_policy=mp_policy,
dp_mesh=device_mesh["dp"])
weight_iterator = safetensors_weights_iterator(weight_dir_list)
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
@@ -146,8 +129,7 @@ def shard_model(
*,
cpu_offload: bool,
reshard_after_forward: bool = True,
mp_policy: Optional[MixedPrecisionPolicy] = None,
dp_mesh: Optional[DeviceMesh] = None,
dp_mesh: DeviceMesh | None = None,
) -> None:
"""
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API.
@@ -174,17 +156,14 @@ def shard_model(
"""
fsdp_kwargs = {
"reshard_after_forward": reshard_after_forward,
"mesh": dp_mesh,
"mp_policy": mp_policy,
"mesh": dp_mesh
}
if cpu_offload:
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
# iterating in reverse to start with
# Shard the model with FSDP, iterating in reverse to start with
# lowest-level modules first
num_layers_sharded = 0
# TODO(will): don't reshard after forward for the last layer to save on the
# all-gather that will immediately happen Shard the model with FSDP,
for n, m in reversed(list(model.named_modules())):
if any([
shard_condition(n, m)
@@ -205,11 +184,11 @@ def shard_model(
# TODO(PY): device mesh for cfg parallel
def load_fsdp_model_from_full_model_state_dict(
model: torch.nn.Module,
full_sd_iterator: Generator[Tuple[str, torch.Tensor], None, None],
full_sd_iterator: Generator[tuple[str, torch.Tensor], None, None],
device: torch.device,
strict: bool = False,
cpu_offload: bool = False,
param_names_mapping: Optional[Callable[[str], tuple[str, Any, Any]]] = None,
param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None = None,
) -> _IncompatibleKeys:
"""
Converting full state dict into a sharded state dict
@@ -231,13 +210,9 @@ def load_fsdp_model_from_full_model_state_dict(
NotImplementedError: If got FSDP with more than 1D.
"""
meta_sharded_sd = model.state_dict()
# s = fully_shard.state(model)
# logger.info(f"type(s): {type(s)}")
# logger.info(f"s: {s}")
# import pdb; pdb.set_trace()
sharded_sd = {}
to_merge_params: DefaultDict[Hashable, Dict[Any, Any]] = defaultdict(dict)
to_merge_params: defaultdict[Hashable, dict[Any, Any]] = defaultdict(dict)
for source_param_name, full_tensor in full_sd_iterator:
assert param_names_mapping is not None
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
+16 -17
View File
@@ -8,8 +8,8 @@ import os
import tempfile
import time
from collections import defaultdict
from collections.abc import Generator
from pathlib import Path
from typing import Generator, List, Optional, Tuple, Union
import filelock
import huggingface_hub.constants
@@ -50,8 +50,7 @@ class DisabledTqdm(tqdm):
super().__init__(*args, **kwargs, disable=True)
def get_lock(model_name_or_path: Union[str, Path],
cache_dir: Optional[str] = None):
def get_lock(model_name_or_path: str | Path, cache_dir: str | None = None):
lock_dir = cache_dir or temp_dir
model_name_or_path = str(model_name_or_path)
os.makedirs(os.path.dirname(lock_dir), exist_ok=True)
@@ -77,10 +76,10 @@ def _shared_pointers(tensors):
def download_weights_from_hf(
model_name_or_path: str,
cache_dir: Optional[str],
allow_patterns: List[str],
revision: Optional[str] = None,
ignore_patterns: Optional[Union[str, List[str]]] = None,
cache_dir: str | None,
allow_patterns: list[str],
revision: str | None = None,
ignore_patterns: str | list[str] | None = None,
) -> str:
"""Download model weights from Hugging Face Hub.
@@ -136,8 +135,8 @@ def download_weights_from_hf(
def download_safetensors_index_file_from_hf(
model_name_or_path: str,
index_file: str,
cache_dir: Optional[str],
revision: Optional[str] = None,
cache_dir: str | None,
revision: str | None = None,
) -> None:
"""Download hf safetensors index file from Hugging Face Hub.
@@ -172,9 +171,9 @@ def download_safetensors_index_file_from_hf(
# Passing both of these to the weight loader functionality breaks.
# So, we use the index_file to
# look up which safetensors files should be used.
def filter_duplicate_safetensors_files(hf_weights_files: List[str],
def filter_duplicate_safetensors_files(hf_weights_files: list[str],
hf_folder: str,
index_file: str) -> List[str]:
index_file: str) -> list[str]:
# model.safetensors.index.json is a mapping from keys in the
# torch state_dict to safetensors file holding that weight.
index_file_name = os.path.join(hf_folder, index_file)
@@ -197,7 +196,7 @@ def filter_duplicate_safetensors_files(hf_weights_files: List[str],
def filter_files_not_needed_for_inference(
hf_weights_files: List[str]) -> List[str]:
hf_weights_files: list[str]) -> list[str]:
"""
Exclude files that are not needed for inference.
@@ -225,8 +224,8 @@ _BAR_FORMAT = "{desc}: {percentage:3.0f}% Completed | {n_fmt}/{total_fmt} [{elap
def safetensors_weights_iterator(
hf_weights_files: List[str]
) -> Generator[Tuple[str, torch.Tensor], None, None]:
hf_weights_files: list[str]
) -> 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
@@ -243,8 +242,8 @@ def safetensors_weights_iterator(
def pt_weights_iterator(
hf_weights_files: List[str]
) -> Generator[Tuple[str, torch.Tensor], None, None]:
hf_weights_files: list[str]
) -> 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
@@ -280,7 +279,7 @@ def default_weight_loader(param: torch.Tensor,
raise
def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> Optional[str]:
def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> str | None:
"""Remap the name of FP8 k/v_scale parameters.
This function handles the remapping of FP8 k/v_scale parameter names.
+12 -13
View File
@@ -1,8 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/parameter.py
from collections.abc import Callable
from fractions import Fraction
from typing import Any, Callable, Tuple, Union
from typing import Any
import torch
from torch.nn import Parameter
@@ -112,9 +113,8 @@ class _ColumnvLLMParameter(BasevLLMParameter):
if shard_offset is None or shard_size is None:
raise ValueError("shard_offset and shard_size must be provided")
if isinstance(
self,
(PackedColumnParameter,
PackedvLLMParameter)) and self.packed_dim == self.output_dim:
self, PackedColumnParameter
| PackedvLLMParameter) and self.packed_dim == self.output_dim:
shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
shard_offset=shard_offset, shard_size=shard_size)
@@ -141,9 +141,8 @@ class _ColumnvLLMParameter(BasevLLMParameter):
assert num_heads is not None
if isinstance(
self,
(PackedColumnParameter,
PackedvLLMParameter)) and self.output_dim == self.packed_dim:
self, PackedColumnParameter
| PackedvLLMParameter) and self.output_dim == self.packed_dim:
shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
shard_offset=shard_offset, shard_size=shard_size)
@@ -230,7 +229,7 @@ class PerTensorScaleParameter(BasevLLMParameter):
self.qkv_idxs = {"q": 0, "k": 1, "v": 2}
super().__init__(**kwargs)
def _shard_id_as_int(self, shard_id: Union[str, int]) -> int:
def _shard_id_as_int(self, shard_id: str | int) -> int:
if isinstance(shard_id, int):
return shard_id
@@ -255,7 +254,7 @@ class PerTensorScaleParameter(BasevLLMParameter):
super().load_row_parallel_weight(*args, **kwargs)
def _load_into_shard_id(self, loaded_weight: torch.Tensor,
shard_id: Union[str, int], **kwargs):
shard_id: str | int, **kwargs):
"""
Slice the parameter data based on the shard id for
loading.
@@ -282,7 +281,7 @@ class PackedColumnParameter(_ColumnvLLMParameter):
for more details on the packed properties.
"""
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
def __init__(self, packed_factor: int | Fraction, packed_dim: int,
**kwargs):
self._packed_factor = packed_factor
self._packed_dim = packed_dim
@@ -297,7 +296,7 @@ class PackedColumnParameter(_ColumnvLLMParameter):
return self._packed_factor
def adjust_shard_indexes_for_packing(self, shard_size,
shard_offset) -> Tuple[Any, Any]:
shard_offset) -> tuple[Any, Any]:
return _adjust_shard_indexes_for_packing(
shard_size=shard_size,
shard_offset=shard_offset,
@@ -315,7 +314,7 @@ class PackedvLLMParameter(ModelWeightParameter):
by accounting for packing and optionally, marlin tile size.
"""
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
def __init__(self, packed_factor: int | Fraction, packed_dim: int,
**kwargs):
self._packed_factor = packed_factor
self._packed_dim = packed_dim
@@ -404,7 +403,7 @@ def permute_param_layout_(param: BasevLLMParameter, input_dim: int,
def _adjust_shard_indexes_for_packing(shard_size, shard_offset,
packed_factor) -> Tuple[Any, Any]:
packed_factor) -> tuple[Any, Any]:
shard_size = shard_size // packed_factor
shard_offset = shard_offset // packed_factor
return shard_size, shard_offset
+23 -23
View File
@@ -8,10 +8,10 @@ import subprocess
import sys
import tempfile
from abc import ABC, abstractmethod
from collections.abc import Callable, Set
from dataclasses import dataclass, field
from functools import lru_cache
from typing import (AbstractSet, Callable, Dict, List, NoReturn, Optional,
Tuple, Type, TypeVar, Union, cast)
from typing import NoReturn, TypeVar, cast
import cloudpickle
from torch import nn
@@ -80,7 +80,7 @@ class _ModelInfo:
architecture: str
@staticmethod
def from_model_cls(model: Type[nn.Module]) -> "_ModelInfo":
def from_model_cls(model: type[nn.Module]) -> "_ModelInfo":
return _ModelInfo(architecture=model.__name__, )
@@ -91,7 +91,7 @@ class _BaseRegisteredModel(ABC):
raise NotImplementedError
@abstractmethod
def load_model_cls(self) -> Type[nn.Module]:
def load_model_cls(self) -> type[nn.Module]:
raise NotImplementedError
@@ -102,10 +102,10 @@ class _RegisteredModel(_BaseRegisteredModel):
"""
interfaces: _ModelInfo
model_cls: Type[nn.Module]
model_cls: type[nn.Module]
@staticmethod
def from_model_cls(model_cls: Type[nn.Module]):
def from_model_cls(model_cls: type[nn.Module]):
return _RegisteredModel(
interfaces=_ModelInfo.from_model_cls(model_cls),
model_cls=model_cls,
@@ -114,7 +114,7 @@ class _RegisteredModel(_BaseRegisteredModel):
def inspect_model_cls(self) -> _ModelInfo:
return self.interfaces
def load_model_cls(self) -> Type[nn.Module]:
def load_model_cls(self) -> type[nn.Module]:
return self.model_cls
@@ -159,16 +159,16 @@ class _LazyRegisteredModel(_BaseRegisteredModel):
return _run_in_subprocess(
lambda: _ModelInfo.from_model_cls(self.load_model_cls()))
def load_model_cls(self) -> Type[nn.Module]:
def load_model_cls(self) -> type[nn.Module]:
mod = importlib.import_module(self.module_name)
return cast(Type[nn.Module], getattr(mod, self.class_name))
return cast(type[nn.Module], getattr(mod, self.class_name))
@lru_cache(maxsize=128)
def _try_load_model_cls(
model_arch: str,
model: _BaseRegisteredModel,
) -> Optional[Type[nn.Module]]:
) -> type[nn.Module] | None:
from fastvideo.v1.platforms import current_platform
current_platform.verify_model_arch(model_arch)
try:
@@ -182,7 +182,7 @@ def _try_load_model_cls(
def _try_inspect_model_cls(
model_arch: str,
model: _BaseRegisteredModel,
) -> Optional[_ModelInfo]:
) -> _ModelInfo | None:
try:
return model.inspect_model_cls()
except Exception:
@@ -194,15 +194,15 @@ def _try_inspect_model_cls(
@dataclass
class _ModelRegistry:
# Keyed by model_arch
models: Dict[str, _BaseRegisteredModel] = field(default_factory=dict)
models: dict[str, _BaseRegisteredModel] = field(default_factory=dict)
def get_supported_archs(self) -> AbstractSet[str]:
def get_supported_archs(self) -> Set[str]:
return self.models.keys()
def register_model(
self,
model_arch: str,
model_cls: Union[Type[nn.Module], str],
model_cls: type[nn.Module] | str,
) -> None:
"""
Register an external model to be used in vLLM.
@@ -232,7 +232,7 @@ class _ModelRegistry:
self.models[model_arch] = model
def _raise_for_unsupported(self, architectures: List[str]) -> NoReturn:
def _raise_for_unsupported(self, architectures: list[str]) -> NoReturn:
all_supported_archs = self.get_supported_archs()
if any(arch in all_supported_archs for arch in architectures):
@@ -244,13 +244,13 @@ class _ModelRegistry:
f"Model architectures {architectures} are not supported for now. "
f"Supported architectures: {all_supported_archs}")
def _try_load_model_cls(self, model_arch: str) -> Optional[Type[nn.Module]]:
def _try_load_model_cls(self, model_arch: str) -> type[nn.Module] | None:
if model_arch not in self.models:
return None
return _try_load_model_cls(model_arch, self.models[model_arch])
def _try_inspect_model_cls(self, model_arch: str) -> Optional[_ModelInfo]:
def _try_inspect_model_cls(self, model_arch: str) -> _ModelInfo | None:
if model_arch not in self.models:
return None
@@ -258,8 +258,8 @@ class _ModelRegistry:
def _normalize_archs(
self,
architectures: Union[str, List[str]],
) -> List[str]:
architectures: str | list[str],
) -> list[str]:
if isinstance(architectures, str):
architectures = [architectures]
if not architectures:
@@ -274,8 +274,8 @@ class _ModelRegistry:
def inspect_model_cls(
self,
architectures: Union[str, List[str]],
) -> Tuple[_ModelInfo, str]:
architectures: str | list[str],
) -> tuple[_ModelInfo, str]:
architectures = self._normalize_archs(architectures)
for arch in architectures:
@@ -287,8 +287,8 @@ class _ModelRegistry:
def resolve_model_cls(
self,
architectures: Union[str, List[str]],
) -> Tuple[Type[nn.Module], str]:
architectures: str | list[str],
) -> tuple[type[nn.Module], str]:
architectures = self._normalize_archs(architectures)
for arch in architectures:
+3 -4
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import Optional, Tuple, Union
import torch
from diffusers.utils import BaseOutput
@@ -32,15 +31,15 @@ class BaseScheduler(ABC):
@abstractmethod
def scale_model_input(self,
sample: torch.Tensor,
timestep: Optional[int] = None) -> torch.Tensor:
timestep: int | None = None) -> torch.Tensor:
pass
@abstractmethod
def step(
self,
model_output: torch.Tensor,
timestep: Union[int, torch.Tensor],
timestep: int | torch.Tensor,
sample: torch.Tensor,
return_dict: bool = True,
) -> Union[BaseOutput, Tuple]:
) -> BaseOutput | tuple:
pass
@@ -20,7 +20,7 @@
# ==============================================================================
from dataclasses import dataclass
from typing import Any, Optional, Tuple, Union
from typing import Any
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
@@ -75,7 +75,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
shift: float = 1.0,
reverse: bool = True,
solver: str = "euler",
n_tokens: Optional[int] = None,
n_tokens: int | None = None,
**kwargs,
):
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
@@ -130,7 +130,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def set_timesteps(
self,
num_inference_steps: int,
device: Union[str, torch.device] = None,
device: str | torch.device = None,
n_tokens: int = 0,
):
"""
@@ -193,7 +193,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def scale_model_input(self,
sample: torch.Tensor,
timestep: Optional[int] = None) -> torch.Tensor:
timestep: int | None = None) -> torch.Tensor:
return sample
def sd3_time_shift(self, t: torch.Tensor):
@@ -202,11 +202,11 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def step(
self,
model_output: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
timestep: float | torch.FloatTensor,
sample: torch.FloatTensor,
return_dict: bool = True,
**kwargs,
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
) -> FlowMatchDiscreteSchedulerOutput | tuple:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
process from the learned model outputs (most often the predicted noise).
@@ -232,7 +232,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
returned, otherwise a tuple is returned where the first element is the sample tensor.
"""
if isinstance(timestep, (int, torch.IntTensor, torch.LongTensor)):
if isinstance(timestep, int | torch.IntTensor | torch.LongTensor):
raise ValueError((
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
@@ -23,7 +23,6 @@
# ==============================================================================
import math
from typing import List, Optional, Tuple, Union
import numpy as np
import torch
@@ -203,7 +202,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
beta_start: float = 0.0001,
beta_end: float = 0.02,
beta_schedule: str = "linear",
trained_betas: Optional[Union[np.ndarray, List[float]]] = None,
trained_betas: np.ndarray | list[float] | None = None,
solver_order: int = 2,
prediction_type: str = "epsilon",
thresholding: bool = False,
@@ -212,16 +211,16 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
predict_x0: bool = True,
solver_type: str = "bh2",
lower_order_final: bool = True,
disable_corrector: Tuple[int, ...] = (),
disable_corrector: tuple[int, ...] = (),
solver_p: SchedulerMixin = None,
use_karras_sigmas: Optional[bool] = False,
use_exponential_sigmas: Optional[bool] = False,
use_beta_sigmas: Optional[bool] = False,
use_flow_sigmas: Optional[bool] = False,
flow_shift: Optional[float] = 1.0,
use_karras_sigmas: bool | None = False,
use_exponential_sigmas: bool | None = False,
use_beta_sigmas: bool | None = False,
use_flow_sigmas: bool | None = False,
flow_shift: float | None = 1.0,
timestep_spacing: str = "linspace",
steps_offset: int = 0,
final_sigmas_type: Optional[str] = "zero", # "zero", "sigma_min"
final_sigmas_type: str | None = "zero", # "zero", "sigma_min"
rescale_betas_zero_snr: bool = False,
):
if self.config.use_beta_sigmas and not is_scipy_available():
@@ -283,21 +282,20 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self.predict_x0 = predict_x0
# setable values
self.num_inference_steps: Optional[int] = None
self.num_inference_steps: int | None = None
timesteps = np.linspace(0,
num_train_timesteps - 1,
num_train_timesteps,
dtype=np.float32)[::-1].copy()
self.timesteps = torch.from_numpy(timesteps)
self.model_outputs = [None] * solver_order
self.timestep_list: List[Union[int,
torch.Tensor]] = [None] * solver_order
self.timestep_list: list[int | torch.Tensor] = [None] * solver_order
self.lower_order_nums = 0
self.disable_corrector = list(disable_corrector)
self.solver_p = solver_p
self.last_sample = None
self._step_index: Optional[int] = None
self._begin_index: Optional[int] = None
self._step_index: int | None = None
self._begin_index: int | None = None
self.sigmas = self.sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
@@ -333,7 +331,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def set_timesteps(self,
num_inference_steps: int,
device: Union[str, torch.device] = None):
device: str | torch.device = None):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
@@ -537,7 +535,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler._sigma_to_alpha_sigma_t
def _sigma_to_alpha_sigma_t(
self, sigma: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
self, sigma: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
if self.config.use_flow_sigmas:
alpha_t = 1 - sigma
sigma_t = sigma
@@ -708,7 +706,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
model_output: torch.Tensor,
*args,
sample: torch.Tensor = None,
order: Optional[int] = None,
order: int | None = None,
**kwargs,
) -> torch.Tensor:
"""
@@ -808,7 +806,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
R_tensor: torch.Tensor = torch.stack(R)
b = torch.tensor(b, device=device)
D1s_tensor: Optional[torch.Tensor] = None
D1s_tensor: torch.Tensor | None = None
if len(D1s) > 0:
D1s_tensor = torch.stack(D1s, dim=1) # (B, K)
# for order 2, we use a simplified version
@@ -842,9 +840,9 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self,
this_model_output: torch.Tensor,
*args,
last_sample: Optional[torch.Tensor] = None,
this_sample: Optional[torch.Tensor] = None,
order: Optional[int] = None,
last_sample: torch.Tensor | None = None,
this_sample: torch.Tensor | None = None,
order: int | None = None,
**kwargs,
) -> torch.Tensor:
"""
@@ -950,7 +948,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
R = torch.stack(R)
b = torch.tensor(b, device=device)
D1s_tensor: Optional[torch.Tensor] = torch.stack(
D1s_tensor: torch.Tensor | None = torch.stack(
D1s, dim=1) if len(D1s) > 0 else None
# for order 1, we use a simplified version
@@ -1016,10 +1014,10 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def step(
self,
model_output: torch.Tensor,
timestep: Union[int, torch.Tensor],
timestep: int | torch.Tensor,
sample: torch.Tensor,
return_dict: bool = True,
) -> Union[SchedulerOutput, Tuple]:
) -> SchedulerOutput | tuple:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with
the multistep UniPC.
+5 -5
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/utils.py
"""Utils for model executor."""
from typing import Any, Dict, List, Optional
from typing import Any
import torch
@@ -58,7 +58,7 @@ def set_random_seed(seed: int) -> None:
def set_weight_attrs(
weight: torch.Tensor,
weight_attrs: Optional[Dict[str, Any]],
weight_attrs: dict[str, Any] | None,
):
"""Set attributes on a weight tensor.
@@ -109,7 +109,7 @@ def extract_layer_index(layer_name: str) -> int:
- "model.encoder.layers.0.sub.1" -> ValueError
"""
subnames = layer_name.split(".")
int_vals: List[int] = []
int_vals: list[int] = []
for subname in subnames:
try:
int_vals.append(int(subname))
@@ -121,8 +121,8 @@ def extract_layer_index(layer_name: str) -> int:
def modulate(x: torch.Tensor,
shift: Optional[torch.Tensor] = None,
scale: Optional[torch.Tensor] = None) -> torch.Tensor:
shift: torch.Tensor | None = None,
scale: torch.Tensor | None = None) -> torch.Tensor:
"""modulate by shift and scale
Args:
+17 -17
View File
@@ -1,8 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from collections.abc import Iterator
from math import prod
from typing import Iterator, Optional, Tuple, Union, cast
from typing import Optional, cast
import numpy as np
import torch
@@ -51,8 +52,8 @@ class ParallelTiledVAE(ABC):
return cast(int, self.config.spatial_compression_ratio)
@property
def scaling_factor(self) -> Union[float, torch.tensor]:
return cast(Union[float, torch.tensor], self.config.scaling_factor)
def scaling_factor(self) -> float | torch.Tensor:
return cast(float | torch.Tensor, self.config.scaling_factor)
@abstractmethod
def _encode(self, *args, **kwargs) -> torch.Tensor:
@@ -160,7 +161,7 @@ class ParallelTiledVAE(ABC):
def _parallel_data_generator(
self, gathered_results,
gathered_dim_metadata) -> Iterator[Tuple[torch.Tensor, int]]:
gathered_dim_metadata) -> Iterator[tuple[torch.Tensor, int]]:
global_idx = 0
for i, per_rank_metadata in enumerate(gathered_dim_metadata):
_start_shape = 0
@@ -412,16 +413,16 @@ class ParallelTiledVAE(ABC):
def enable_tiling(
self,
tile_sample_min_height: Optional[int] = None,
tile_sample_min_width: Optional[int] = None,
tile_sample_min_num_frames: Optional[int] = None,
tile_sample_stride_height: Optional[int] = None,
tile_sample_stride_width: Optional[int] = None,
tile_sample_stride_num_frames: Optional[int] = None,
blend_num_frames: Optional[int] = None,
use_tiling: Optional[bool] = None,
use_temporal_tiling: Optional[bool] = None,
use_parallel_tiling: Optional[bool] = None,
tile_sample_min_height: int | None = None,
tile_sample_min_width: int | None = None,
tile_sample_min_num_frames: int | None = None,
tile_sample_stride_height: int | None = None,
tile_sample_stride_width: int | None = None,
tile_sample_stride_num_frames: int | None = None,
blend_num_frames: int | None = None,
use_tiling: bool | None = None,
use_temporal_tiling: bool | None = None,
use_parallel_tiling: bool | None = None,
) -> None:
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
@@ -485,8 +486,7 @@ class DiagonalGaussianDistribution:
device=self.parameters.device,
dtype=self.parameters.dtype)
def sample(self,
generator: Optional[torch.Generator] = None) -> torch.Tensor:
def sample(self, generator: torch.Generator | None = None) -> torch.Tensor:
# make sure sample is on the same device as the parameters and has same dtype
sample = randn_tensor(
self.mean.shape,
@@ -517,7 +517,7 @@ class DiagonalGaussianDistribution:
def nll(
self, sample: torch.Tensor,
dims: Tuple[int, ...] = (1, 2, 3)) -> torch.Tensor:
dims: tuple[int, ...] = (1, 2, 3)) -> torch.Tensor:
if self.deterministic:
return torch.Tensor([0.0])
logtwopi = np.log(2.0 * np.pi)
+25 -23
View File
@@ -15,8 +15,6 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Optional, Tuple, Union
import numpy as np
import torch
import torch.nn as nn
@@ -32,7 +30,7 @@ def prepare_causal_attention_mask(
height_width: int,
dtype: torch.dtype,
device: torch.device,
batch_size: Optional[int] = None) -> torch.Tensor:
batch_size: int | None = None) -> torch.Tensor:
indices = torch.arange(1, num_frames + 1, dtype=torch.int32, device=device)
indices_blocks = indices.repeat_interleave(height_width)
x, y = torch.meshgrid(indices_blocks, indices_blocks, indexing="xy")
@@ -72,7 +70,7 @@ class HunyuanVAEAttention(nn.Module):
def forward(self,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
attention_mask: torch.Tensor | None = None) -> torch.Tensor:
residual = hidden_states
batch_size, sequence_length, _ = hidden_states.shape
@@ -121,10 +119,10 @@ class HunyuanVideoCausalConv3d(nn.Module):
self,
in_channels: int,
out_channels: int,
kernel_size: Union[int, Tuple[int, int, int]] = 3,
stride: Union[int, Tuple[int, int, int]] = 1,
padding: Union[int, Tuple[int, int, int]] = 0,
dilation: Union[int, Tuple[int, int, int]] = 1,
kernel_size: int | tuple[int, int, int] = 3,
stride: int | tuple[int, int, int] = 1,
padding: int | tuple[int, int, int] = 0,
dilation: int | tuple[int, int, int] = 1,
bias: bool = True,
pad_mode: str = "replicate",
) -> None:
@@ -163,11 +161,11 @@ class HunyuanVideoUpsampleCausal3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: Optional[int] = None,
out_channels: int | None = None,
kernel_size: int = 3,
stride: int = 1,
bias: bool = True,
upsample_factor: Tuple[int, ...] = (2, 2, 2),
upsample_factor: tuple[int, ...] = (2, 2, 2),
) -> None:
super().__init__()
@@ -213,7 +211,7 @@ class HunyuanVideoDownsampleCausal3D(nn.Module):
def __init__(
self,
channels: int,
out_channels: Optional[int] = None,
out_channels: int | None = None,
padding: int = 1,
kernel_size: int = 3,
bias: bool = True,
@@ -239,7 +237,7 @@ class HunyuanVideoResnetBlockCausal3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: Optional[int] = None,
out_channels: int | None = None,
dropout: float = 0.0,
groups: int = 32,
eps: float = 1e-6,
@@ -313,7 +311,7 @@ class HunyuanVideoMidBlock3D(nn.Module):
non_linearity=resnet_act_fn,
)
]
attentions: list[Optional[HunyuanVAEAttention]] = []
attentions: list[HunyuanVAEAttention | None] = []
for _ in range(num_layers):
if self.add_attention:
@@ -349,7 +347,9 @@ class HunyuanVideoMidBlock3D(nn.Module):
hidden_states = self._gradient_checkpointing_func(
self.resnets[0], hidden_states)
for attn, resnet in zip(self.attentions, self.resnets[1:]):
for attn, resnet in zip(self.attentions,
self.resnets[1:],
strict=False):
if attn is not None:
batch_size, num_channels, num_frames, height, width = hidden_states.shape
hidden_states = hidden_states.permute(0, 2, 3, 4,
@@ -371,7 +371,9 @@ class HunyuanVideoMidBlock3D(nn.Module):
else:
hidden_states = self.resnets[0](hidden_states)
for attn, resnet in zip(self.attentions, self.resnets[1:]):
for attn, resnet in zip(self.attentions,
self.resnets[1:],
strict=False):
if attn is not None:
batch_size, num_channels, num_frames, height, width = hidden_states.shape
hidden_states = hidden_states.permute(0, 2, 3, 4,
@@ -404,7 +406,7 @@ class HunyuanVideoDownBlock3D(nn.Module):
resnet_act_fn: str = "silu",
resnet_groups: int = 32,
add_downsample: bool = True,
downsample_stride: Tuple[int, ...] | int = 2,
downsample_stride: tuple[int, ...] | int = 2,
downsample_padding: int = 1,
) -> None:
super().__init__()
@@ -466,7 +468,7 @@ class HunyuanVideoUpBlock3D(nn.Module):
resnet_act_fn: str = "silu",
resnet_groups: int = 32,
add_upsample: bool = True,
upsample_scale_factor: Tuple[int, ...] = (2, 2, 2),
upsample_scale_factor: tuple[int, ...] = (2, 2, 2),
) -> None:
super().__init__()
resnets = []
@@ -525,13 +527,13 @@ class HunyuanVideoEncoder3D(nn.Module):
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: Tuple[str, ...] = (
down_block_types: tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
),
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
block_out_channels: tuple[int, ...] = (128, 256, 512, 512),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
@@ -546,7 +548,7 @@ class HunyuanVideoEncoder3D(nn.Module):
block_out_channels[0],
kernel_size=3,
stride=1)
self.mid_block: Optional[HunyuanVideoMidBlock3D] = None
self.mid_block: HunyuanVideoMidBlock3D | None = None
self.down_blocks = nn.ModuleList([])
output_channel = block_out_channels[0]
@@ -649,13 +651,13 @@ class HunyuanVideoDecoder3D(nn.Module):
self,
in_channels: int = 3,
out_channels: int = 3,
up_block_types: Tuple[str, ...] = (
up_block_types: tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
),
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
block_out_channels: tuple[int, ...] = (128, 256, 512, 512),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
@@ -832,7 +834,7 @@ class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
self,
sample: torch.Tensor,
sample_posterior: bool = False,
generator: Optional[torch.Generator] = None,
generator: torch.Generator | None = None,
) -> torch.Tensor:
r"""
Args:
+5 -5
View File
@@ -10,7 +10,7 @@
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
from typing import Any, List, Optional, Tuple
from typing import Any
import torch
from einops import rearrange
@@ -102,7 +102,7 @@ def base_conv3d(x,
return out
def cal_outsize(input_sizes, kernel_sizes, stride, padding) -> List:
def cal_outsize(input_sizes, kernel_sizes, stride, padding) -> list:
stride_d, stride_h, stride_w = stride
padding_d, padding_h, padding_w = padding
dilation_d, dilation_h, dilation_w = 1, 1, 1
@@ -445,8 +445,8 @@ def base_group_norm_with_zero_pad(x,
class CausalConvChannelLast(CausalConv):
time_causal_padding: Tuple[Any, ...]
time_uncausal_padding: Tuple[Any, ...]
time_causal_padding: tuple[Any, ...]
time_uncausal_padding: tuple[Any, ...]
def __init__(self, chan_in, chan_out, kernel_size, **kwargs) -> None:
super().__init__(chan_in, chan_out, kernel_size, **kwargs)
@@ -1121,7 +1121,7 @@ class AutoencoderKLStepvideo(nn.Module, ParallelTiledVAE):
self,
sample: torch.Tensor,
sample_posterior: bool = False,
generator: Optional[torch.Generator] = None,
generator: torch.Generator | None = None,
) -> torch.Tensor:
"""
Args:
+13 -11
View File
@@ -16,7 +16,6 @@
import contextvars
from contextlib import contextmanager
from typing import Optional, Tuple, Union
import torch
import torch.nn as nn
@@ -68,9 +67,9 @@ class WanCausalConv3d(nn.Conv3d):
self,
in_channels: int,
out_channels: int,
kernel_size: Union[int, Tuple[int, int, int]],
stride: Union[int, Tuple[int, int, int]] = 1,
padding: Union[int, Tuple[int, int, int]] = 0,
kernel_size: int | tuple[int, int, int],
stride: int | tuple[int, int, int] = 1,
padding: int | tuple[int, int, int] = 0,
) -> None:
super().__init__(
in_channels=in_channels,
@@ -79,9 +78,9 @@ class WanCausalConv3d(nn.Conv3d):
stride=stride,
padding=padding,
)
self.padding: Tuple[int, int, int]
self.padding: tuple[int, int, int]
# Set up causal padding
self._padding: Tuple[int, ...] = (self.padding[2], self.padding[2],
self._padding: tuple[int, ...] = (self.padding[2], self.padding[2],
self.padding[1], self.padding[1],
2 * self.padding[0], 0)
self.padding = (0, 0, 0)
@@ -434,7 +433,8 @@ class WanMidBlock(nn.Module):
x = self.resnets[0](x)
# Process through attention and residual blocks
for attn, resnet in zip(self.attentions, self.resnets[1:]):
for attn, resnet in zip(self.attentions, self.resnets[1:],
strict=False):
if attn is not None:
x = attn(x)
@@ -488,7 +488,8 @@ class WanEncoder3d(nn.Module):
# downsample blocks
self.down_blocks = nn.ModuleList([])
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
for i, (in_dim,
out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=False)):
# residual (+attention) blocks
for _ in range(num_res_blocks):
self.down_blocks.append(
@@ -589,7 +590,7 @@ class WanUpBlock(nn.Module):
out_dim: int,
num_res_blocks: int,
dropout: float = 0.0,
upsample_mode: Optional[str] = None,
upsample_mode: str | None = None,
non_linearity: str = "silu",
):
super().__init__()
@@ -687,7 +688,8 @@ class WanDecoder3d(nn.Module):
# upsample blocks
self.up_blocks = nn.ModuleList([])
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
for i, (in_dim,
out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=False)):
# residual (+attention) blocks
if i > 0:
in_dim = in_dim // 2
@@ -944,7 +946,7 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
self,
sample: torch.Tensor,
sample_posterior: bool = False,
generator: Optional[torch.Generator] = None,
generator: torch.Generator | None = None,
) -> torch.Tensor:
"""
Args:
+11 -15
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import os
from typing import Callable, List, Optional, Tuple, Union
from collections.abc import Callable
import numpy as np
import PIL.Image
@@ -29,8 +29,7 @@ else:
}
def pil_to_numpy(
images: Union[List[PIL.Image.Image], PIL.Image.Image]) -> np.ndarray:
def pil_to_numpy(images: list[PIL.Image.Image] | PIL.Image.Image) -> np.ndarray:
r"""
Convert a PIL image or a list of PIL images to NumPy arrays.
@@ -69,9 +68,7 @@ def numpy_to_pt(images: np.ndarray) -> torch.Tensor:
return images
def normalize(
images: Union[np.ndarray,
torch.Tensor]) -> Union[np.ndarray, torch.Tensor]:
def normalize(images: np.ndarray | torch.Tensor) -> np.ndarray | torch.Tensor:
r"""
Normalize an image array to [-1,1].
@@ -87,9 +84,8 @@ def normalize(
def load_image(
image: Union[str, PIL.Image.Image],
convert_method: Optional[Callable[[PIL.Image.Image],
PIL.Image.Image]] = None
image: str | PIL.Image.Image,
convert_method: Callable[[PIL.Image.Image], PIL.Image.Image] | None = None
) -> PIL.Image.Image:
"""
Loads `image` to a PIL Image.
@@ -132,11 +128,11 @@ def load_image(
def get_default_height_width(
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
image: PIL.Image.Image | np.ndarray | torch.Tensor,
vae_scale_factor: int,
height: Optional[int] = None,
width: Optional[int] = None,
) -> Tuple[int, int]:
height: int | None = None,
width: int | None = None,
) -> tuple[int, int]:
r"""
Returns the height and width of the image, downscaled to the next integer multiple of `vae_scale_factor`.
@@ -179,12 +175,12 @@ def get_default_height_width(
def resize(
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
image: PIL.Image.Image | np.ndarray | torch.Tensor,
height: int,
width: int,
resize_mode: str = "default", # "default", "fill", "crop"
resample: str = "lanczos",
) -> Union[PIL.Image.Image, np.ndarray, torch.Tensor]:
) -> PIL.Image.Image | np.ndarray | torch.Tensor:
"""
Resize image.
+22 -166
View File
@@ -5,25 +5,19 @@ Base class for composed pipelines.
This module defines the base class for pipelines that are composed of multiple stages.
"""
import argparse
import os
from abc import ABC, abstractmethod
from copy import deepcopy
from typing import Any, Dict, List, Optional, Union, cast
from typing import Any, cast
import torch
from fastvideo.v1.configs.pipelines import (PipelineConfig,
get_pipeline_config_cls_for_name)
from fastvideo.v1.distributed import (init_distributed_environment,
initialize_model_parallel,
model_parallel_is_initialized)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import PipelineStage
from fastvideo.v1.utils import (maybe_download_model, shallow_asdict,
from fastvideo.v1.utils import (maybe_download_model,
verify_model_config_and_directory)
logger = init_logger(__name__)
@@ -39,25 +33,20 @@ class ComposedPipelineBase(ABC):
"""
is_video_pipeline: bool = False # To be overridden by video pipelines
_required_config_modules: List[str] = []
_required_config_modules: list[str] = []
# TODO(will): args should support both inference args and training args
def __init__(self,
model_path: str,
fastvideo_args: FastVideoArgs,
config: Optional[Dict[str, Any]] = None,
required_config_modules: Optional[List[str]] = None):
config: dict[str, Any] | None = None):
"""
Initialize the pipeline. After __init__, the pipeline should be ready to
use. The pipeline should be stateless and not hold any batch state.
"""
self.fastvideo_args = fastvideo_args
self.model_path = model_path
self._stages: List[PipelineStage] = []
self._stage_name_mapping: Dict[str, PipelineStage] = {}
if required_config_modules is not None:
self._required_config_modules = required_config_modules
self._stages: list[PipelineStage] = []
self._stage_name_mapping: dict[str, PipelineStage] = {}
if self._required_config_modules is None:
raise NotImplementedError(
@@ -70,150 +59,31 @@ class ComposedPipelineBase(ABC):
else:
self.config = config
self.maybe_init_distributed_environment(fastvideo_args)
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
self.modules = self.load_modules(fastvideo_args)
if fastvideo_args.training_mode:
if fastvideo_args.log_validation:
self.initialize_validation_pipeline(fastvideo_args)
self.initialize_training_pipeline(fastvideo_args)
self.initialize_pipeline(fastvideo_args)
# logger.info("Creating pipeline stages...")
# self.create_pipeline_stages(fastvideo_args)
logger.info("Creating pipeline stages...")
self.create_pipeline_stages(fastvideo_args)
if fastvideo_args.training_mode:
logger.info("Creating training pipeline stages...")
self.create_training_stages(fastvideo_args)
else:
logger.info("Creating pipeline stages...")
self.create_pipeline_stages(fastvideo_args)
def initialize_training_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"if training_mode is True, the pipeline must implement this method")
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"if log_validation is True, the pipeline must implement this method"
)
@classmethod
def from_pretrained(cls,
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[
Union[str
| PipelineConfig]] = None,
args: Optional[argparse.Namespace] = None,
required_config_modules: Optional[List[str]] = None,
**kwargs) -> "ComposedPipelineBase":
config = None
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, PipelineConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = kwargs
else:
config_args = shallow_asdict(config)
config_args.update(kwargs)
if args.inference_mode:
fastvideo_args = FastVideoArgs(model_path=model_path,
device_str=device or "cuda" if
torch.cuda.is_available() else "cpu",
**config_args)
fastvideo_args.model_path = model_path
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
) else "cpu"
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
else:
assert args is not None, "args must be provided for training mode"
fastvideo_args = TrainingArgs.from_cli_args(args)
# TODO(will): fix this so that its not so ugly
fastvideo_args.model_path = model_path
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
) else "cpu"
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
# we use cpu offload for training
fastvideo_args.use_cpu_offload = False
# make sure we are in training mode
fastvideo_args.inference_mode = False
# we hijack the precision to be the master weight type so that the
# model is loaded with the correct precision. Subsequently we will
# use FSDP2's MixedPrecisionPolicy to set the precision for the
# fwd, bwd, and other operations' precision.
fastvideo_args.precision = fastvideo_args.master_weight_type
assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
fastvideo_args.check_fastvideo_args()
logger.info(f"fastvideo_args in from_pretrained: {fastvideo_args}")
return cls(model_path,
fastvideo_args,
required_config_modules=required_config_modules)
def maybe_init_distributed_environment(self, fastvideo_args: FastVideoArgs):
if model_parallel_is_initialized():
return
local_rank = int(os.environ.get("LOCAL_RANK", -1))
world_size = int(os.environ.get("WORLD_SIZE", -1))
rank = int(os.environ.get("RANK", -1))
if local_rank == -1 or world_size == -1 or rank == -1:
raise ValueError(
"Local rank, world size, and rank must be set. Use torchrun to launch the script."
)
torch.cuda.set_device(local_rank)
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
initialize_model_parallel(
tensor_model_parallel_size=fastvideo_args.tp_size,
sequence_model_parallel_size=fastvideo_args.sp_size)
device = torch.device(f"cuda:{local_rank}")
fastvideo_args.device = device
def get_module(self, module_name: str, default_value: Any = None) -> Any:
if module_name not in self.modules:
return default_value
def get_module(self, module_name: str) -> Any:
return self.modules[module_name]
def add_module(self, module_name: str, module: Any):
self.modules[module_name] = module
def _load_config(self, model_path: str) -> Dict[str, Any]:
def _load_config(self, model_path: str) -> dict[str, Any]:
model_path = maybe_download_model(self.model_path)
self.model_path = model_path
# fastvideo_args.downloaded_model_path = model_path
logger.info("Model path: %s", model_path)
config = verify_model_config_and_directory(model_path)
return cast(Dict[str, Any], config)
return cast(dict[str, Any], config)
@property
def required_config_modules(self) -> List[str]:
def required_config_modules(self) -> list[str]:
"""
List of modules that are required by the pipeline. The names should match
the diffusers directory and model_index.json file. These modules will be
@@ -231,7 +101,7 @@ class ComposedPipelineBase(ABC):
return self._required_config_modules
@property
def stages(self) -> List[PipelineStage]:
def stages(self) -> list[PipelineStage]:
"""
List of stages in the pipeline.
"""
@@ -244,26 +114,13 @@ class ComposedPipelineBase(ABC):
"""
raise NotImplementedError
# @abstractmethod
# def create_validation_stages(self, fastvideo_args: FastVideoArgs):
# """
# Create the validation pipeline stages.
# """
# raise NotImplementedError
def create_training_stages(self, fastvideo_args: FastVideoArgs):
"""
Create the training pipeline stages.
"""
raise NotImplementedError
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
return
def load_modules(self, fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
def load_modules(self, fastvideo_args: FastVideoArgs) -> dict[str, Any]:
"""
Load the modules from the config.
"""
@@ -279,21 +136,19 @@ class ComposedPipelineBase(ABC):
modules_config
) > 1, "model_index.json must contain at least one pipeline module"
for module_name in self.required_config_modules:
required_modules = [
"vae", "text_encoder", "transformer", "scheduler", "tokenizer"
]
for module_name in required_modules:
if module_name not in modules_config:
raise ValueError(
f"model_index.json must contain a {module_name} module")
logger.info("Diffusers config passed sanity checks")
# all the component models used by the pipeline
required_modules = self.required_config_modules
logger.info("Loading required modules: %s", required_modules)
modules = {}
for module_name, (transformers_or_diffusers,
architecture) in modules_config.items():
if module_name not in required_modules:
logger.info("Skipping module %s", module_name)
continue
component_model_path = os.path.join(self.model_path, module_name)
module = PipelineComponentLoader.load_module(
module_name=module_name,
@@ -309,6 +164,7 @@ class ComposedPipelineBase(ABC):
logger.warning("Overwriting module %s", module_name)
modules[module_name] = module
required_modules = self.required_config_modules
# Check if all required modules were loaded
for module_name in required_modules:
if module_name not in modules or modules[module_name] is None:
@@ -342,7 +198,7 @@ class ComposedPipelineBase(ABC):
# Execute each stage
logger.info("Running pipeline stages: %s",
self._stage_name_mapping.keys())
# logger.info("Batch: %s", batch)
logger.info("Batch: %s", batch)
for stage in self.stages:
batch = stage(batch, fastvideo_args)
+35 -35
View File
@@ -8,7 +8,7 @@ in a functional manner, reducing the need for explicit parameter passing.
"""
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Union
from typing import Any
import torch
@@ -30,81 +30,81 @@ class ForwardBatch:
# specific arguments.
data_type: str
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None
generator: torch.Generator | list[torch.Generator] | None = None
# Image inputs
image_path: Optional[str] = None
image_embeds: List[torch.Tensor] = field(default_factory=list)
image_path: str | None = None
image_embeds: list[torch.Tensor] = field(default_factory=list)
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
negative_prompt: Optional[Union[str, List[str]]] = None
prompt_path: Optional[str] = None
prompt: str | list[str] | None = None
negative_prompt: str | list[str] | None = None
prompt_path: str | None = None
output_path: str = "outputs/"
# Primary encoder embeddings
prompt_embeds: List[torch.Tensor] = field(default_factory=list)
negative_prompt_embeds: Optional[List[torch.Tensor]] = None
prompt_attention_mask: Optional[List[torch.Tensor]] = None
negative_attention_mask: Optional[List[torch.Tensor]] = None
clip_embedding_pos: Optional[List[torch.Tensor]] = None
clip_embedding_neg: Optional[List[torch.Tensor]] = None
prompt_embeds: list[torch.Tensor] = field(default_factory=list)
negative_prompt_embeds: list[torch.Tensor] | None = None
prompt_attention_mask: list[torch.Tensor] | None = None
negative_attention_mask: list[torch.Tensor] | None = None
clip_embedding_pos: list[torch.Tensor] | None = None
clip_embedding_neg: list[torch.Tensor] | None = None
# Additional text-related parameters
max_sequence_length: Optional[int] = None
prompt_template: Optional[Dict[str, Any]] = None
max_sequence_length: int | None = None
prompt_template: dict[str, Any] | None = None
do_classifier_free_guidance: bool = False
# Batch info
batch_size: Optional[int] = None
batch_size: int | None = None
num_videos_per_prompt: int = 1
seed: Optional[int] = None
seeds: Optional[List[int]] = None
seed: int | None = None
seeds: list[int] | None = None
# Tracking if embeddings are already processed
is_prompt_processed: bool = False
# Latent tensors
latents: Optional[torch.Tensor] = None
noise_pred: Optional[torch.Tensor] = None
image_latent: Optional[torch.Tensor] = None
latents: torch.Tensor | None = None
noise_pred: torch.Tensor | None = None
image_latent: torch.Tensor | None = None
# Latent dimensions
height_latents: Optional[int] = None
width_latents: Optional[int] = None
height_latents: int | None = None
width_latents: int | None = None
num_frames: int = 1 # Default for image models
num_frames_round_down: bool = False # Whether to round down num_frames if it's not divisible by num_gpus
# Original dimensions (before VAE scaling)
height: Optional[int] = None
width: Optional[int] = None
fps: Optional[int] = None
height: int | None = None
width: int | None = None
fps: int | None = None
# Timesteps
timesteps: Optional[torch.Tensor] = None
timestep: Optional[Union[torch.Tensor, float, int]] = None
step_index: Optional[int] = None
timesteps: torch.Tensor | None = None
timestep: torch.Tensor | float | int | None = None
step_index: int | None = None
# Scheduler parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
eta: float = 0.0
sigmas: Optional[List[float]] = None
sigmas: list[float] | None = None
n_tokens: Optional[int] = None
n_tokens: int | None = None
# Other parameters that may be needed by specific schedulers
extra_step_kwargs: Dict[str, Any] = field(default_factory=dict)
extra_step_kwargs: dict[str, Any] = field(default_factory=dict)
# Component modules (populated by the pipeline)
modules: Dict[str, Any] = field(default_factory=dict)
modules: dict[str, Any] = field(default_factory=dict)
# Final output (after pipeline completion)
output: Any = None
# Extra parameters that might be needed by specific pipeline implementations
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
# Misc
save_video: bool = True
@@ -112,7 +112,7 @@ class ForwardBatch:
# TeaCache parameters
enable_teacache: bool = False
teacache_params: Optional[TeaCacheParams | WanTeaCacheParams] = None
teacache_params: TeaCacheParams | WanTeaCacheParams | None = None
def __post_init__(self):
"""Initialize dependent fields after dataclass initialization."""
+6 -6
View File
@@ -4,9 +4,9 @@
import importlib
import pkgutil
from collections.abc import Set
from dataclasses import dataclass, field
from functools import lru_cache
from typing import AbstractSet, Dict, Optional, Tuple, Type
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
@@ -17,14 +17,14 @@ logger = init_logger(__name__)
@dataclass
class _PipelineRegistry:
# Keyed by pipeline_arch
pipelines: Dict[str, Optional[Type[ComposedPipelineBase]]] = field(
default_factory=dict)
pipelines: dict[str, type[ComposedPipelineBase]
| None] = field(default_factory=dict)
def get_supported_archs(self) -> AbstractSet[str]:
def get_supported_archs(self) -> Set[str]:
return self.pipelines.keys()
def _try_load_pipeline_cls(
self, pipeline_arch: str) -> Optional[Type[ComposedPipelineBase]]:
self, pipeline_arch: str) -> type[ComposedPipelineBase] | None:
if pipeline_arch not in self.pipelines:
return None
@@ -33,7 +33,7 @@ class _PipelineRegistry:
def resolve_pipeline_cls(
self,
architecture: str,
) -> Tuple[Type[ComposedPipelineBase], str]:
) -> tuple[type[ComposedPipelineBase], str]:
if not architecture:
logger.warning("No pipeline architecture is specified")
@@ -1,563 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
T2V Data Preprocessing pipeline implementation.
This module contains an implementation of the T2V Data Preprocessing pipeline
using the modular pipeline architecture.
"""
import gc
import multiprocessing
import os
from concurrent.futures import ProcessPoolExecutor
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.v1.dataset import getdataset
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import TextEncodingStage
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class PreprocessPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
args,
):
# Initialize class variables for data sharing
self.video_data = {} # Store video metadata and paths
self.latent_data = {} # Store latent tensors
self.preprocess_validation_text(fastvideo_args, args)
self.preprocess_video_and_text(fastvideo_args, args)
def preprocess_video_and_text(self, fastvideo_args: FastVideoArgs, args):
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(combined_parquet_dir, exist_ok=True)
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
# Get how many samples have already been processed
start_idx = 0
for root, _, files in os.walk(combined_parquet_dir):
for file in files:
if file.endswith('.parquet'):
table = pq.read_table(os.path.join(root, file))
start_idx += table.num_rows
# Loading dataset
train_dataset = getdataset(args, start_idx=start_idx)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=False)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
num_processed_samples = 0
# Add progress bar for video preprocessing
pbar = tqdm(train_dataloader,
desc="Processing videos",
unit="batch",
disable=local_rank != 0)
for batch_idx, data in enumerate(pbar):
if data is None:
continue
with torch.inference_mode():
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, pixel_values in enumerate(data["pixel_values"]):
if not torch.all(
pixel_values == 0): # Check if all values are zero
valid_indices.append(i)
num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples
valid_data = {
"pixel_values":
torch.stack(
[data["pixel_values"][i] for i in valid_indices]),
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
"fps": [data["fps"][i] for i in valid_indices],
"duration": [data["duration"][i] for i in valid_indices],
}
# VAE
with torch.autocast("cuda", dtype=torch.float32):
latents = self.get_module("vae").encode(
valid_data["pixel_values"].to(
fastvideo_args.device)).mean
# Get corresponding captions for this batch
batch_captions = valid_data["text"]
batch = ForwardBatch(
data_type="video",
prompt=batch_captions,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
prompt_embeds, prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
assert prompt_embeds.shape[0] == prompt_attention_mask.shape[0]
# Get sequence lengths from attention masks (number of 1s)
seq_lens = prompt_attention_mask.sum(dim=1)
# Create a list to store non-padded embeddings and masks
non_padded_embeds = []
non_padded_masks = []
# Process each item in the batch
for i in range(prompt_embeds.size(0)):
seq_len = seq_lens[i].item()
# Slice the embeddings and masks to keep only non-padding parts
non_padded_embeds.append(prompt_embeds[i, :seq_len])
non_padded_masks.append(prompt_attention_mask[i, :seq_len])
# Update the tensors with non-padded versions
prompt_embeds = non_padded_embeds
prompt_attention_mask = non_padded_masks
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
# Get the corresponding latent and info using video name
latent = latents[idx].cpu()
video_name = os.path.basename(video_path).split(".")[0]
height, width = valid_data["pixel_values"][idx].shape[-2:]
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
).astype(np.uint8)
# Create record for Parquet dataset
record = {
"id": video_name,
"vae_latent_bytes": vae_latent.tobytes(),
"vae_latent_shape": list(vae_latent.shape),
"vae_latent_dtype": str(vae_latent.dtype),
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape":
list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx],
"media_type": "video",
"width": width,
"height": height,
"num_frames": latents[idx].shape[1],
"duration_sec": float(valid_data["duration"][idx]),
"fps": float(valid_data["fps"][idx]),
}
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = [
pa.array([record["id"] for record in batch_data]),
pa.array(
[record["vae_latent_bytes"] for record in batch_data],
type=pa.binary()),
pa.array(
[record["vae_latent_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array(
[record["vae_latent_dtype"] for record in batch_data]),
pa.array([
record["text_embedding_bytes"] for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_embedding_shape"] for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_embedding_dtype"] for record in batch_data
]),
pa.array([
record["text_attention_mask_bytes"]
for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_attention_mask_shape"]
for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_attention_mask_dtype"]
for record in batch_data
]),
pa.array([record["file_name"] for record in batch_data]),
pa.array([record["caption"] for record in batch_data]),
pa.array([record["media_type"] for record in batch_data]),
pa.array([record["width"] for record in batch_data],
type=pa.int32()),
pa.array([record["height"] for record in batch_data],
type=pa.int32()),
pa.array([record["num_frames"] for record in batch_data],
type=pa.int32()),
pa.array([record["duration_sec"] for record in batch_data],
type=pa.float32()),
pa.array([record["fps"] for record in batch_data],
type=pa.float32()),
]
table = pa.Table.from_arrays(
arrays, names=[f.name for f in pyarrow_schema])
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info(f"Collected batch with {len(table)} samples")
if num_processed_samples >= args.flush_frequency:
assert hasattr(self, 'all_tables') and self.all_tables
print(f"Combining {len(self.all_tables)} batches...")
combined_table = pa.concat_tables(self.all_tables)
assert len(combined_table) == num_processed_samples
print(f"Total samples collected: {len(combined_table)}")
# Calculate total number of chunks needed, discarding remainder
total_chunks = max(
num_processed_samples // args.samples_per_file, 1)
print(
f"Fixed samples per parquet file: {args.samples_per_file}")
print(f"Total number of parquet files: {total_chunks}")
print(
f"Total samples to be processed: {total_chunks * args.samples_per_file} (discarding {num_processed_samples % args.samples_per_file} samples)"
)
# Split work among processes
num_workers = int(min(multiprocessing.cpu_count(),
total_chunks))
chunks_per_worker = (total_chunks + num_workers -
1) // num_workers
print(
f"Using {num_workers} workers to process {total_chunks} chunks"
)
logger.info(f"Chunks per worker: {chunks_per_worker}")
# Prepare work ranges
work_ranges = []
for i in range(num_workers):
start_idx = i * chunks_per_worker
end_idx = min((i + 1) * chunks_per_worker, total_chunks)
if start_idx < total_chunks:
work_ranges.append(
(start_idx, end_idx, combined_table, i,
combined_parquet_dir, args.samples_per_file))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
for work_range in work_ranges
}
for future in tqdm(futures, desc="Processing chunks"):
try:
written = future.result()
total_written += written
logger.info(
f"Processed chunk with {written} samples")
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]}: {str(e)}"
)
# Retry failed ranges sequentially
if failed_ranges:
logger.warning(
f"Retrying {len(failed_ranges)} failed ranges sequentially"
)
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(
work_range)
except Exception as e:
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]} after retry: {str(e)}"
)
logger.info(f"Total samples written: {total_written}")
num_processed_samples = 0
self.all_tables = []
def preprocess_validation_text(self, fastvideo_args: FastVideoArgs, args):
# Create Parquet dataset directory for validation
validation_parquet_dir = os.path.join(args.output_dir,
"validation_parquet_dataset")
os.makedirs(validation_parquet_dir, exist_ok=True)
with open(args.validation_prompt_txt, encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for validation text preprocessing
pbar = tqdm(enumerate(prompts),
desc="Processing validation prompts",
unit="prompt")
for prompt_idx, prompt in pbar:
with torch.inference_mode():
# Text Encoder
batch = ForwardBatch(
data_type="video",
prompt=prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
prompt_embeds = result_batch.prompt_embeds[0]
prompt_attention_mask = result_batch.prompt_attention_mask[0]
file_name = prompt.split(".")[0]
# Get the sequence length from attention mask (number of 1s)
seq_len = prompt_attention_mask.sum().item()
# Slice the embeddings to keep only the non-padding parts
text_embedding = prompt_embeds[0, :seq_len].cpu().numpy()
text_attention_mask = prompt_attention_mask[
0, :seq_len].cpu().numpy().astype(np.uint8)
# Log the shapes after removing padding
logger.info(
f"Shape after removing padding - Embeddings: {text_embedding.shape}, Mask: {text_attention_mask.shape}"
)
# Create record for Parquet dataset
record = {
"id": file_name,
"vae_latent_bytes": b"", # Not available for validation
"vae_latent_shape": [],
"vae_latent_dtype": "",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape": list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": file_name,
"caption": prompt,
"media_type": "video",
"width": 0, # Not available for validation
"height": 0, # Not available for validation
"num_frames": 0, # Not available for validation
"duration_sec": 0.0, # Not available for validation
"fps": 0.0, # Not available for validation
}
batch_data.append(record)
logger.info(f"Saved validation sample: {file_name}")
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = [
pa.array([record["id"] for record in batch_data]),
pa.array([record["vae_latent_bytes"] for record in batch_data],
type=pa.binary()),
pa.array([record["vae_latent_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array([record["vae_latent_dtype"] for record in batch_data]),
pa.array(
[record["text_embedding_bytes"] for record in batch_data],
type=pa.binary()),
pa.array(
[record["text_embedding_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array(
[record["text_embedding_dtype"] for record in batch_data]),
pa.array([
record["text_attention_mask_bytes"] for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_attention_mask_shape"] for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_attention_mask_dtype"] for record in batch_data
]),
pa.array([record["file_name"] for record in batch_data]),
pa.array([record["caption"] for record in batch_data]),
pa.array([record["media_type"] for record in batch_data]),
pa.array([record["width"] for record in batch_data],
type=pa.int32()),
pa.array([record["height"] for record in batch_data],
type=pa.int32()),
pa.array([record["num_frames"] for record in batch_data],
type=pa.int32()),
pa.array([record["duration_sec"] for record in batch_data],
type=pa.float32()),
pa.array([record["fps"] for record in batch_data],
type=pa.float32()),
]
table = pa.Table.from_arrays(arrays,
names=[f.name for f in pyarrow_schema])
write_pbar.update(1)
write_pbar.close()
logger.info(f"Total validation samples: {len(table)}")
work_range = (0, 1, table, 0, validation_parquet_dir, len(table))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=1) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
}
for future in tqdm(futures, desc="Processing chunks"):
try:
total_written += future.result()
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]}: {str(e)}"
)
# Retry failed ranges sequentially
if failed_ranges:
logger.warning(
f"Retrying {len(failed_ranges)} failed ranges sequentially")
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(work_range)
except Exception as e:
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]} after retry: {str(e)}"
)
logger.info(f"Total validation samples written: {total_written}")
# Clear memory
del table
gc.collect() # Force garbage collection
@staticmethod
def process_chunk_range(args):
start_idx, end_idx, table, worker_id, output_dir, samples_per_file = args
try:
total_written = 0
num_samples = len(table)
# Create worker-specific subdirectory
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
os.makedirs(worker_dir, exist_ok=True)
# Check how many files there are already in the dir, and update i accordingly
num_parquets = 0
for root, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
for i in range(start_idx, end_idx):
start_sample = i * samples_per_file
end_sample = min((i + 1) * samples_per_file, num_samples)
chunk = table.slice(start_sample, end_sample - start_sample)
# Create chunk file in worker's directory
chunk_path = os.path.join(
worker_dir, f"data_chunk_{i + num_parquets}.parquet")
temp_path = chunk_path + '.tmp'
try:
# Write to temporary file
pq.write_table(chunk, temp_path, compression='zstd')
# Rename temporary file to final file
if os.path.exists(chunk_path):
os.remove(
chunk_path) # Remove existing file if it exists
os.rename(temp_path, chunk_path)
total_written += len(chunk)
except Exception as e:
# Clean up temporary file if it exists
if os.path.exists(temp_path):
os.remove(temp_path)
raise e
return total_written
except Exception as e:
logger.error(
f"Error processing chunks {start_idx}-{end_idx} for worker {worker_id}: {str(e)}"
)
raise
EntryClass = PreprocessPipeline
+8 -9
View File
@@ -5,7 +5,8 @@ Denoising stage for diffusion pipelines.
import importlib.util
import inspect
from typing import Any, Dict, Iterable, List, Optional
from collections.abc import Iterable
from typing import Any
import torch
from einops import rearrange
@@ -74,8 +75,7 @@ class DenoisingStage(PipelineStage):
)
# Setup precision and autocast settings
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
target_dtype = torch.bfloat16
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
@@ -84,7 +84,6 @@ class DenoisingStage(PipelineStage):
), get_sequence_model_parallel_rank()
sp_group = world_size > 1
if sp_group:
# b c t h w -> b t n s h w
latents = rearrange(batch.latents,
"b t (n s) h w -> b t n s h w",
n=world_size).contiguous()
@@ -110,7 +109,7 @@ class DenoisingStage(PipelineStage):
def dict_to_3d_list(mask_strategy,
t_max=50,
l_max=60,
h_max=24) -> List:
h_max=24) -> list:
result = [[[None for _ in range(h_max)] for _ in range(l_max)]
for _ in range(t_max)]
if mask_strategy is None:
@@ -190,7 +189,7 @@ class DenoisingStage(PipelineStage):
# Predict noise residual
with torch.autocast(device_type="cuda",
dtype=torch.bfloat16,
dtype=target_dtype,
enabled=autocast_enabled):
# TODO(will-refactor): all of this should be in the stage's init
@@ -296,7 +295,7 @@ class DenoisingStage(PipelineStage):
return batch
def prepare_extra_func_kwargs(self, func, kwargs) -> Dict[str, Any]:
def prepare_extra_func_kwargs(self, func, kwargs) -> dict[str, Any]:
"""
Prepare extra kwargs for the scheduler step / denoise step.
@@ -315,8 +314,8 @@ class DenoisingStage(PipelineStage):
return extra_step_kwargs
def progress_bar(self,
iterable: Optional[Iterable] = None,
total: Optional[int] = None) -> tqdm:
iterable: Iterable | None = None,
total: int | None = None) -> tqdm:
"""
Create a progress bar for the denoising process.
+3 -4
View File
@@ -2,7 +2,6 @@
"""
Encoding stage for diffusion pipelines.
"""
from typing import Optional
import PIL.Image
import torch
@@ -137,7 +136,7 @@ class EncodingStage(PipelineStage):
def retrieve_latents(self,
encoder_output: torch.Tensor,
generator: Optional[torch.Generator] = None,
generator: torch.Generator | None = None,
sample_mode: str = "sample"):
if sample_mode == "sample":
return encoder_output.sample(generator)
@@ -151,8 +150,8 @@ class EncodingStage(PipelineStage):
self,
image: PIL.Image.Image,
vae_scale_factor: int,
height: Optional[int] = None,
width: Optional[int] = None,
height: int | None = None,
width: int | None = None,
resize_mode: str = "default", # "default", "fill", "crop"
) -> torch.Tensor:
image = [image]
+8 -16
View File
@@ -56,22 +56,19 @@ class TextEncodingStage(PipelineStage):
fastvideo_args.text_encoder_configs)
for tokenizer, text_encoder, encoder_config, preprocess_func, postprocess_func in zip(
self.tokenizers, self.text_encoders,
self.tokenizers,
self.text_encoders,
fastvideo_args.text_encoder_configs,
fastvideo_args.preprocess_text_funcs,
fastvideo_args.postprocess_text_funcs):
fastvideo_args.postprocess_text_funcs,
strict=False):
if fastvideo_args.use_cpu_offload:
text_encoder = text_encoder.to(fastvideo_args.device)
assert isinstance(batch.prompt, (str, list))
if isinstance(batch.prompt, str):
batch.prompt = [batch.prompt]
texts = []
for prompt_str in batch.prompt:
texts.append(preprocess_func(prompt_str))
text_inputs = tokenizer(texts,
**encoder_config.tokenizer_kwargs).to(
fastvideo_args.device)
assert isinstance(batch.prompt, str)
text = preprocess_func(batch.prompt)
text_inputs = tokenizer(text, **encoder_config.tokenizer_kwargs).to(
fastvideo_args.device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
with set_forward_context(current_timestep=0, attn_metadata=None):
@@ -83,8 +80,6 @@ class TextEncodingStage(PipelineStage):
prompt_embeds = postprocess_func(outputs)
batch.prompt_embeds.append(prompt_embeds)
if batch.prompt_attention_mask is not None:
batch.prompt_attention_mask.append(attention_mask)
if batch.do_classifier_free_guidance:
assert isinstance(batch.negative_prompt, str)
@@ -105,9 +100,6 @@ class TextEncodingStage(PipelineStage):
assert batch.negative_prompt_embeds is not None
batch.negative_prompt_embeds.append(negative_prompt_embeds)
if batch.negative_attention_mask is not None:
batch.negative_attention_mask.append(
negative_attention_mask)
if fastvideo_args.use_cpu_offload:
text_encoder.to('cpu')
@@ -9,7 +9,7 @@ using the modular pipeline architecture.
import os
from copy import deepcopy
from typing import Any, Dict
from typing import Any
import torch
from huggingface_hub import hf_hub_download
@@ -95,7 +95,7 @@ class StepVideoPipeline(ComposedPipelineBase):
))
torch.ops.load_library(lib_path)
def load_modules(self, fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
def load_modules(self, fastvideo_args: FastVideoArgs) -> dict[str, Any]:
"""
Load the modules from the config.
"""
-841
View File
@@ -1,841 +0,0 @@
import gc
import os
import sys
import time
import traceback
from abc import ABC, abstractmethod
from collections import deque
from copy import deepcopy
import imageio
import numpy as np
import torch
import torchvision
from diffusers.optimization import get_scheduler
from einops import rearrange
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
# import torch.distributed as dist
import wandb
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.parquet_datasets import ParquetVideoTextDataset
from fastvideo.v1.distributed import cleanup_dist_env_and_memory, get_sp_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
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.training_utils import (
_clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas, save_checkpoint)
from fastvideo.v1.pipelines.wan.wan_pipeline import WanValidationPipeline
logger = init_logger(__name__)
# Manual gradient checking flag - set to True to enable gradient verification
ENABLE_GRADIENT_CHECK = False
GRADIENT_CHECK_DTYPE = torch.bfloat16
class TrainingPipeline(ComposedPipelineBase, ABC):
"""
A pipeline for training a model. All training pipelines should inherit from this class.
All reusable components and code should be implemented in this class.
"""
_required_config_modules = ["scheduler", "transformer"]
def initialize_training_pipeline(self, fastvideo_args: TrainingArgs):
logger.info("Initializing training pipeline...")
self.device = fastvideo_args.device
self.sp_group = get_sp_group()
self.world_size = self.sp_group.world_size
self.rank = self.sp_group.rank
self.local_rank = self.sp_group.local_rank
self.transformer = self.get_module("transformer")
assert self.transformer is not None
self.transformer.requires_grad_(True)
self.transformer.train()
args = fastvideo_args
noise_scheduler = self.modules["scheduler"]
params_to_optimize = self.transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
init_steps = 0
logger.info("optimizer: %s", optimizer)
# todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * self.world_size,
num_training_steps=args.max_train_steps * self.world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = ParquetVideoTextDataset(
args.data_path,
batch_size=args.train_batch_size,
rank=self.rank,
world_size=self.world_size,
cfg_rate=args.cfg,
num_latent_t=args.num_latent_t)
train_dataloader = StatefulDataLoader(
train_dataset,
batch_size=args.train_batch_size,
num_workers=args.
dataloader_num_workers, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=True)
self.lr_scheduler = lr_scheduler
self.train_dataset = train_dataset
self.train_dataloader = train_dataloader
self.init_steps = init_steps
self.optimizer = optimizer
self.noise_scheduler = noise_scheduler
# self.noise_random_generator = noise_random_generator
# num_update_steps_per_epoch = math.ceil(
# len(train_dataloader) / args.gradient_accumulation_steps *
# args.sp_size / args.train_sp_batch_size)
# args.num_train_epochs = math.ceil(args.max_train_steps /
# num_update_steps_per_epoch)
if self.rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
@abstractmethod
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"Training pipelines must implement this method")
@abstractmethod
def train_one_step(self, transformer, model_type, optimizer, lr_scheduler,
loader, noise_scheduler, noise_random_generator,
gradient_accumulation_steps, sp_size,
precondition_outputs, max_grad_norm, weighting_scheme,
logit_mean, logit_std, mode_scale):
"""
Train one step of the model.
"""
raise NotImplementedError(
"Training pipeline must implement this method")
def log_validation(self, transformer, fastvideo_args, global_step):
fastvideo_args.inference_mode = True
fastvideo_args.use_cpu_offload = False
if not fastvideo_args.log_validation:
return
if self.validation_pipeline is None:
raise ValueError("Validation pipeline is not set")
# Create sampling parameters if not provided
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
# Prepare validation prompts
print('fastvideo_args.validation_prompt_dir',
fastvideo_args.validation_prompt_dir)
validation_dataset = ParquetVideoTextDataset(
fastvideo_args.validation_prompt_dir,
batch_size=1,
rank=0,
world_size=1,
cfg_rate=0,
num_latent_t=args.num_latent_t)
validation_dataloader = StatefulDataLoader(
validation_dataset,
batch_size=1,
num_workers=1, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=False)
transformer.requires_grad_(False)
for p in transformer.parameters():
p.requires_grad = False
transformer.eval()
# Add the transformer to the validation pipeline
self.validation_pipeline.add_module("transformer", transformer)
self.validation_pipeline.latent_preparation_stage.transformer = transformer
self.validation_pipeline.denoising_stage.transformer = transformer
# Process each validation prompt
videos = []
captions = []
for _, embeddings, masks, infos in validation_dataloader:
logger.info(f"infos: {infos}")
caption = infos['caption']
captions.append(caption)
prompt_embeds = embeddings.to(fastvideo_args.device).to(torch.bfloat16)
prompt_attention_mask = masks.to(fastvideo_args.device).to(torch.bfloat16)
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
logger.info('embed dtype', prompt_embeds.dtype)
num_frames = (fastvideo_args.num_latent_t - 1) * 4 + 1
logger.info(f"validation num_frames: {num_frames}")
# Prepare batch for validation
# print('shape of embeddings', prompt_embeds.shape)
batch = ForwardBatch(
# **shallow_asdict(sampling_param),
data_type="video",
latents=None,
# seed=sampling_param.seed,
# data_type="video",
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
# make sure we use the same height, width, and num_frames as the training pipeline
height=args.num_height,
width=args.num_width,
num_frames=num_frames,
# num_inference_steps=fastvideo_args.validation_sampling_steps,
num_inference_steps=50,
# guidance_scale=fastvideo_args.validation_guidance_scale,
guidance_scale=1,
n_tokens=n_tokens,
do_classifier_free_guidance=False,
eta=0.0,
extra={},
)
# Run validation inference
with torch.autocast("cuda", dtype=torch.bfloat16):
with torch.inference_mode():
output_batch = self.validation_pipeline.forward(
batch, fastvideo_args)
samples = output_batch.output
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
videos.append(frames)
# Log validation results
rank = int(os.environ.get("RANK", 0))
if rank == 0:
video_filenames = []
video_captions = []
for i, video in enumerate(videos):
caption = captions[i]
os.makedirs(fastvideo_args.output_dir, exist_ok=True)
filename = os.path.join(
fastvideo_args.output_dir,
f"validation_step_{global_step}_video_{i}.mp4")
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
video_captions.append(
caption) # Store the caption for each video
logs = {
"validation_videos": [
wandb.Video(filename,
caption=caption) for filename, caption in zip(
video_filenames, video_captions)
]
}
wandb.log(logs, step=global_step)
# Re-enable gradients for training
transformer.requires_grad_(True)
transformer.train()
gc.collect()
torch.cuda.empty_cache()
def gradient_check_parameters(self,
transformer,
latents,
encoder_hidden_states,
encoder_attention_mask,
timesteps,
target,
eps=5e-2,
max_params_to_check=2000):
"""
Verify gradients using finite differences for FSDP models with GRADIENT_CHECK_DTYPE.
Uses standard tolerances for GRADIENT_CHECK_DTYPE precision.
"""
# Move all inputs to CPU and clear GPU memory
inputs_cpu = {
'latents': latents.cpu(),
'encoder_hidden_states': encoder_hidden_states.cpu(),
'encoder_attention_mask': encoder_attention_mask.cpu(),
'timesteps': timesteps.cpu(),
'target': target.cpu()
}
del latents, encoder_hidden_states, encoder_attention_mask, timesteps, target
torch.cuda.empty_cache()
def compute_loss():
# Move inputs to GPU, compute loss, cleanup
inputs_gpu = {
k:
v.to(self.fastvideo_args.device,
dtype=GRADIENT_CHECK_DTYPE
if k != 'encoder_attention_mask' else None)
for k, v in inputs_cpu.items()
}
# Use GRADIENT_CHECK_DTYPE for more accurate gradient checking
# with torch.autocast(enabled=False, device_type="cuda"):
with torch.autocast("cuda", dtype=GRADIENT_CHECK_DTYPE):
with set_forward_context(
current_timestep=inputs_gpu['timesteps'],
attn_metadata=None):
model_pred = transformer(
hidden_states=inputs_gpu['latents'],
encoder_hidden_states=inputs_gpu[
'encoder_hidden_states'],
timestep=inputs_gpu['timesteps'],
encoder_attention_mask=inputs_gpu[
'encoder_attention_mask'],
return_dict=False)[0]
if self.fastvideo_args.precondition_outputs:
sigmas = get_sigmas(self.noise_scheduler,
inputs_gpu['latents'].device,
inputs_gpu['timesteps'],
n_dim=inputs_gpu['latents'].ndim,
dtype=inputs_gpu['latents'].dtype)
model_pred = inputs_gpu['latents'] - model_pred * sigmas
target_adjusted = inputs_gpu['target']
else:
target_adjusted = inputs_gpu['target']
loss = torch.mean((model_pred - target_adjusted)**2)
# Cleanup and return
loss_cpu = loss.cpu()
del inputs_gpu, model_pred, target_adjusted
if 'sigmas' in locals(): del sigmas
torch.cuda.empty_cache()
return loss_cpu.to(self.fastvideo_args.device)
try:
# Get analytical gradients
transformer.zero_grad()
analytical_loss = compute_loss()
analytical_loss.backward()
# Check gradients for selected parameters
absolute_errors = []
param_count = 0
rank = int(os.environ.get("RANK", 0))
sp_group = get_sp_group()
for name, param in transformer.named_parameters():
sp_group.barrier()
# skip scale_shift_table because it is not sharded
if 'scale_shift_table' in name:
continue
if isinstance(param.grad, torch.distributed.tensor.DTensor):
l = param.grad.full_tensor()
distributed = True
else:
l = param.grad
distributed = False
continue
# logger.info(f"rank: {rank}, name: {name}, param: {param.shape}, grad: {param.grad.shape}, distributed: {distributed}", local_main_process_only=False)
# logger.info(f"rank: {rank}, name: {name}, type of param: {type(param)}, type of grad: {type(param.grad)}", local_main_process_only=False)
# logger.info(f"rank: {rank}, name: {name}, param: {param}, grad: {param.grad}", local_main_process_only=False)
if not (param.requires_grad and param.grad is not None
and param_count < max_params_to_check
and l.abs().max() > 5e-4):
continue
if not distributed:
if rank != 0:
continue
# Get local parameter and gradient tensors
local_param = param._local_tensor if hasattr(
param, '_local_tensor') else param
local_grad = param.grad._local_tensor if hasattr(
param.grad, '_local_tensor') else param.grad
# logger.info(f"rank: {rank}, local_param: {local_param.shape}, local_grad: {local_grad.shape}", local_main_process_only=False)
# Find first significant gradient element
flat_param = local_param.data.view(-1)
flat_grad = local_grad.view(-1)
# logger.info(f"rank: {rank}, flat_param: {flat_param.shape}, flat_grad: {flat_grad.shape}", local_main_process_only=False)
check_idx = next((i for i in range(min(10, flat_param.numel()))
if abs(flat_grad[i]) > 1e-4), 0)
# logger.info(f"rank: {rank}, check_idx: {check_idx}", local_main_process_only=False)
# Store original values
orig_value = flat_param[check_idx].item()
analytical_grad = flat_grad[check_idx].item()
# Compute numerical gradient
for delta in [eps, -eps]:
with torch.no_grad():
# only have a single rank modify the parameter
if rank == 0:
flat_param[check_idx] = orig_value + delta
loss = compute_loss()
if delta > 0: loss_plus = loss.item()
else: loss_minus = loss.item()
# Restore parameter and compute error
with torch.no_grad():
flat_param[check_idx] = orig_value
numerical_grad = (loss_plus - loss_minus) / (2 * eps)
abs_error = abs(analytical_grad - numerical_grad)
rel_error = abs_error / max(abs(analytical_grad),
abs(numerical_grad), 1e-3)
absolute_errors.append(abs_error)
if self.rank <= 0:
logger.info(
f"{name}[{check_idx}]: analytical={analytical_grad:.6f}, "
f"numerical={numerical_grad:.6f}, abs_error={abs_error:.2e}, rel_error={rel_error:.2%}"
)
# param_count += 1
# Compute and log statistics
if self.rank <= 0:
if absolute_errors:
min_err, max_err, mean_err = min(absolute_errors), max(
absolute_errors
), sum(absolute_errors) / len(absolute_errors)
logger.info(
f"Gradient check stats: min={min_err:.2e}, max={max_err:.2e}, mean={mean_err:.2e}"
)
wandb.log({
"grad_check/min_abs_error":
min_err,
"grad_check/max_abs_error":
max_err,
"grad_check/mean_abs_error":
mean_err,
"grad_check/analytical_loss":
analytical_loss.item(),
})
return max_err
return float('inf')
except Exception as e:
logger.error(f"Gradient check failed: {e}")
traceback.print_exc()
return float('inf')
def setup_gradient_check(self, args, loader_iter, noise_scheduler,
noise_random_generator):
"""
Setup and perform gradient check on a fresh batch.
Args:
args: Training arguments
loader_iter: Data loader iterator
noise_scheduler: Noise scheduler for diffusion
noise_random_generator: Random number generator for noise
Returns:
float or None: Maximum gradient error or None if check is disabled/fails
"""
if not ENABLE_GRADIENT_CHECK:
return None
try:
# Get a fresh batch and process it exactly like train_one_step
check_latents, check_encoder_hidden_states, check_encoder_attention_mask, check_infos = next(
loader_iter)
# Process exactly like in train_one_step but use GRADIENT_CHECK_DTYPE
check_latents = check_latents.to(self.fastvideo_args.device,
dtype=GRADIENT_CHECK_DTYPE)
check_encoder_hidden_states = check_encoder_hidden_states.to(
self.fastvideo_args.device, dtype=GRADIENT_CHECK_DTYPE)
check_latents = normalize_dit_input("wan", check_latents)
batch_size = check_latents.shape[0]
check_noise = torch.randn_like(check_latents)
check_u = compute_density_for_timestep_sampling(
weighting_scheme=args.weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=args.logit_mean,
logit_std=args.logit_std,
mode_scale=args.mode_scale,
)
check_indices = (check_u *
noise_scheduler.config.num_train_timesteps).long()
check_timesteps = noise_scheduler.timesteps[check_indices].to(
device=check_latents.device)
check_sigmas = get_sigmas(
noise_scheduler,
check_latents.device,
check_timesteps,
n_dim=check_latents.ndim,
dtype=check_latents.dtype,
)
check_noisy_model_input = (
1.0 - check_sigmas) * check_latents + check_sigmas * check_noise
# Compute target exactly like train_one_step
if args.precondition_outputs:
check_target = check_latents
else:
check_target = check_noise - check_latents
# Perform gradient check with the exact same inputs as training
max_grad_error = self.gradient_check_parameters(
transformer=self.transformer,
latents=
check_noisy_model_input, # Use noisy input like in training
encoder_hidden_states=check_encoder_hidden_states,
encoder_attention_mask=check_encoder_attention_mask,
timesteps=check_timesteps,
target=check_target,
max_params_to_check=100 # Check more parameters
)
if max_grad_error > 5e-2:
logger.error(
f"❌ Large gradient error detected: {max_grad_error:.2e}")
else:
logger.info(
f"✅ Gradient check passed: max error {max_grad_error:.2e}")
return max_grad_error
except Exception as e:
logger.error(f"Gradient check setup failed: {e}")
traceback.print_exc()
return None
class WanTrainingPipeline(TrainingPipeline):
"""
A training pipeline for Wan.
"""
_required_config_modules = ["scheduler", "transformer"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
pass
def create_training_stages(self, fastvideo_args: FastVideoArgs):
pass
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
self.latents = None
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(fastvideo_args)
args_copy.inference_mode = True
args_copy.vae_config.load_encoder = False
# TODO(will): clean this up
args_copy.precision = "bf16"
validation_pipeline = WanValidationPipeline.from_pretrained(
args.model_path, args=args_copy)
self.validation_pipeline = validation_pipeline
def train_one_step(
self,
transformer,
model_type,
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
precondition_outputs,
max_grad_norm,
weighting_scheme,
logit_mean,
logit_std,
mode_scale,
):
self.modules["transformer"].requires_grad_(True)
self.modules["transformer"].train()
total_loss = 0.0
optimizer.zero_grad()
for _ in range(gradient_accumulation_steps):
logger.info(f"Rank {self.rank}: Training step {_}", local_main_process_only=False)
if self.latents is None:
(
self.latents,
self.encoder_hidden_states,
self.encoder_attention_mask,
self.infos,
) = next(loader_iter)
latents = self.latents
encoder_hidden_states = self.encoder_hidden_states
encoder_attention_mask = self.encoder_attention_mask
infos = self.infos
logger.info(f"Rank {self.rank}: Training step {_} loaded data", local_main_process_only=False)
latents = latents.to(self.fastvideo_args.device,
dtype=torch.bfloat16)
encoder_hidden_states = encoder_hidden_states.to(
self.fastvideo_args.device, dtype=torch.bfloat16)
latents = normalize_dit_input(model_type, latents)
batch_size = latents.shape[0]
noise = torch.randn_like(latents)
u = compute_density_for_timestep_sampling(
weighting_scheme=weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=logit_mean,
logit_std=logit_std,
mode_scale=mode_scale,
)
indices = (u * noise_scheduler.config.num_train_timesteps).long()
timesteps = noise_scheduler.timesteps[indices].to(
device=latents.device)
if sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
sp_group.broadcast(timesteps, src=0)
sigmas = get_sigmas(
noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
with torch.autocast("cuda", dtype=torch.bfloat16):
input_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if 'hunyuan' in model_type:
input_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
with set_forward_context(current_timestep=timesteps,
attn_metadata=None):
model_pred = transformer(**input_kwargs)[0]
if precondition_outputs:
model_pred = noisy_model_input - model_pred * sigmas
if precondition_outputs:
target = latents
else:
target = noise - latents
loss = (torch.mean((model_pred.float() - target.float())**2) /
gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
sp_group = get_sp_group()
sp_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
total_loss += avg_loss.item()
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,
)
optimizer.step()
print('device after optimizer step',
next(transformer.named_parameters())[1].device)
lr_scheduler.step()
print('device after scheduler step',
next(transformer.named_parameters())[1].device)
return total_loss, grad_norm.item()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
):
args = fastvideo_args
self.fastvideo_args = args
train_dataloader = self.train_dataloader
init_steps = self.init_steps
lr_scheduler = self.lr_scheduler
optimizer = self.optimizer
noise_scheduler = self.noise_scheduler
noise_random_generator = None
from diffusers import FlowMatchEulerDiscreteScheduler
noise_scheduler = FlowMatchEulerDiscreteScheduler()
# Train!
total_batch_size = (self.world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
logger.info("***** Running training *****")
# logger.info(f" Num examples = {len(train_dataset)}")
# logger.info(f" Dataloader size = {len(train_dataloader)}")
# logger.info(f" Num Epochs = {args.num_train_epochs}")
logger.info(f" Resume training from step {init_steps}")
logger.info(
f" Instantaneous batch size per device = {args.train_batch_size}")
logger.info(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
logger.info(
f" Gradient Accumulation steps = {args.gradient_accumulation_steps}"
)
logger.info(f" Total optimization steps = {args.max_train_steps}")
logger.info(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in self.transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
logger.info(
f" Master weight dtype: {self.transformer.parameters().__next__().dtype}"
)
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=self.local_rank > 0,
)
loader_iter = iter(train_dataloader)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader_iter)
# get gpu memory usage
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info(
f"GPU memory usage before train_one_step: {gpu_memory_usage} MB")
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.perf_counter()
loss, grad_norm = self.train_one_step(
self.transformer,
# args.model_type,
"wan",
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.precondition_outputs,
args.max_grad_norm,
args.weighting_scheme,
args.logit_mean,
args.logit_std,
args.mode_scale,
)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info(
f"GPU memory usage after train_one_step: {gpu_memory_usage} MB")
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
# Manual gradient checking - only at first step
if step == 1 and ENABLE_GRADIENT_CHECK:
logger.info(f"Performing gradient check at step {step}")
self.setup_gradient_check(args, loader_iter, noise_scheduler,
noise_random_generator)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
if self.rank <= 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
},
step=step,
)
if step % args.checkpointing_steps == 0:
# Your existing checkpoint saving code
save_checkpoint(self.transformer, self.rank, args.output_dir,
step)
self.transformer.train()
self.sp_group.barrier()
if args.log_validation and step % args.validation_steps == 0:
self.log_validation(self.transformer, args, step)
save_checkpoint(self.transformer, self.rank, args.output_dir,
args.max_train_steps)
if get_sp_group():
cleanup_dist_env_and_memory()
def main(args):
logger.info("Starting training pipeline...")
pipeline = WanTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.fastvideo_args
pipeline.forward(None, args)
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
print(args)
main(args)
-310
View File
@@ -1,310 +0,0 @@
import json
import math
import os
from typing import List, Optional, Tuple, Union
import torch
import torch.distributed as dist
import torch.distributed.tensor
from torch.distributed.fsdp import FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import StateDictType
from fastvideo.v1.logger import init_logger
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = False
logger = init_logger(__name__)
def compute_density_for_timestep_sampling(
weighting_scheme: str,
batch_size: int,
generator,
logit_mean: float = None,
logit_std: float = None,
mode_scale: float = None,
):
"""
Compute the density for sampling the timesteps when doing SD3 training.
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
"""
if weighting_scheme == "logit_normal":
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
u = torch.normal(
mean=logit_mean,
std=logit_std,
size=(batch_size, ),
device="cpu",
generator=generator,
)
u = torch.nn.functional.sigmoid(u)
elif weighting_scheme == "mode":
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2)**2 - 1 + u)
else:
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
return u
def get_sigmas(noise_scheduler,
device,
timesteps,
n_dim=4,
dtype=torch.float32):
sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype)
schedule_timesteps = noise_scheduler.timesteps.to(device)
timesteps = timesteps.to(device)
step_indices = [(schedule_timesteps == t).nonzero().item()
for t in timesteps]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < n_dim:
sigma = sigma.unsqueeze(-1)
return sigma
def save_checkpoint(transformer, rank, output_dir, step):
# Configure FSDP to save full state dict
FSDP.set_state_dict_type(
transformer,
state_dict_type=StateDictType.FULL_STATE_DICT,
state_dict_config=FullStateDictConfig(offload_to_cpu=True,
rank0_only=True),
)
# Now get the state dict
cpu_state = transformer.state_dict()
# Save it (only on rank 0 since we used rank0_only=True)
if rank <= 0:
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.pt")
torch.save(cpu_state, weight_path)
config_dict = transformer.hf_config
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
config_path = os.path.join(save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
logger.info("--> checkpoint saved at step %s to %s", step, weight_path)
def _clip_grad_norm_while_handling_failing_dtensor_cases(
parameters: Union[torch.Tensor, List[torch.Tensor]],
max_norm: float,
norm_type: float = 2.0,
error_if_nonfinite: bool = False,
foreach: Optional[bool] = None,
pp_mesh: Optional[torch.distributed.device_mesh.DeviceMesh] = None,
) -> Optional[torch.Tensor]:
global _HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES
if not _HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES:
try:
return clip_grad_norm_(parameters, max_norm, norm_type,
error_if_nonfinite, foreach, pp_mesh)
except NotImplementedError as e:
if "DTensor does not support cross-mesh operation" in str(e):
# https://github.com/pytorch/pytorch/issues/134212
logger.warning(
"DTensor does not support cross-mesh operation. If you haven't fully tensor-parallelized your "
"model, while combining other parallelisms such as FSDP, it could be the reason for this error. "
"Gradient clipping will be skipped and gradient norm will not be logged."
)
except Exception as e:
logger.warning(
f"An error occurred while clipping gradients: {e}. Gradient clipping will be skipped and gradient "
f"norm will not be logged.")
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = True
return None
# Copied from https://github.com/pytorch/torchtitan/blob/4a169701555ab9bd6ca3769f9650ae3386b84c6e/torchtitan/utils.py#L362
@torch.no_grad()
def clip_grad_norm_(
parameters: Union[torch.Tensor, List[torch.Tensor]],
max_norm: float,
norm_type: float = 2.0,
error_if_nonfinite: bool = False,
foreach: Optional[bool] = None,
pp_mesh: Optional[torch.distributed.device_mesh.DeviceMesh] = None,
) -> torch.Tensor:
r"""
Clip the gradient norm of parameters.
Gradient norm clipping requires computing the gradient norm over the entire model.
`torch.nn.utils.clip_grad_norm_` only computes gradient norm along DP/FSDP/TP dimensions.
We need to manually reduce the gradient norm across PP stages.
See https://github.com/pytorch/torchtitan/issues/596 for details.
Args:
parameters (`torch.Tensor` or `List[torch.Tensor]`):
Tensors that will have gradients normalized.
max_norm (`float`):
Maximum norm of the gradients after clipping.
norm_type (`float`, defaults to `2.0`):
Type of p-norm to use. Can be `inf` for infinity norm.
error_if_nonfinite (`bool`, defaults to `False`):
If `True`, an error is thrown if the total norm of the gradients from `parameters` is `nan`, `inf`, or `-inf`.
foreach (`bool`, defaults to `None`):
Use the faster foreach-based implementation. If `None`, use the foreach implementation for CUDA and CPU native tensors
and silently fall back to the slow implementation for other device types.
pp_mesh (`torch.distributed.device_mesh.DeviceMesh`, defaults to `None`):
Pipeline parallel device mesh. If not `None`, will reduce gradient norm across PP stages.
Returns:
`torch.Tensor`:
Total norm of the gradients
"""
grads = [p.grad for p in parameters if p.grad is not None]
# TODO(aryan): Wait for next Pytorch release to use `torch.nn.utils.get_total_norm`
# total_norm = torch.nn.utils.get_total_norm(grads, norm_type, error_if_nonfinite, foreach)
total_norm = _get_total_norm(grads, norm_type, error_if_nonfinite, foreach)
# If total_norm is a DTensor, the placements must be `torch.distributed._tensor.ops.math_ops._NormPartial`.
# We can simply reduce the DTensor to get the total norm in this tensor's process group
# and then convert it to a local tensor.
# It has two purposes:
# 1. to make sure the total norm is computed correctly when PP is used (see below)
# 2. to return a reduced total_norm tensor whose .item() would return the correct value
if isinstance(total_norm, torch.distributed.tensor.DTensor):
# Will reach here if any non-PP parallelism is used.
# If only using PP, total_norm will be a local tensor.
total_norm = total_norm.full_tensor()
if pp_mesh is not None:
if math.isinf(norm_type):
dist.all_reduce(total_norm,
op=dist.ReduceOp.MAX,
group=pp_mesh.get_group())
else:
total_norm **= norm_type
dist.all_reduce(total_norm,
op=dist.ReduceOp.SUM,
group=pp_mesh.get_group())
total_norm **= 1.0 / norm_type
_clip_grads_with_norm_(parameters, max_norm, total_norm, foreach)
return total_norm
@torch.no_grad()
def _clip_grads_with_norm_(
parameters: Union[torch.Tensor, List[torch.Tensor]],
max_norm: float,
total_norm: torch.Tensor,
foreach: Optional[bool] = None,
) -> None:
if isinstance(parameters, torch.Tensor):
parameters = [parameters]
grads = [p.grad for p in parameters if p.grad is not None]
max_norm = float(max_norm)
if len(grads) == 0:
return
grouped_grads: dict[Tuple[torch.device, torch.dtype],
Tuple[List[List[torch.Tensor]],
List[int]]] = (_group_tensors_by_device_and_dtype(
[grads])) # type: ignore[assignment]
clip_coef = max_norm / (total_norm + 1e-6)
# Note: multiplying by the clamped coef is redundant when the coef is clamped to 1, but doing so
# avoids a `if clip_coef < 1:` conditional which can require a CPU <=> device synchronization
# when the gradients do not reside in CPU memory.
clip_coef_clamped = torch.clamp(clip_coef, max=1.0)
for (device, _), ([device_grads], _) in grouped_grads.items():
if (foreach is None and _has_foreach_support(device_grads, device)) or (
foreach and _device_has_foreach_support(device)):
torch._foreach_mul_(device_grads, clip_coef_clamped.to(device))
elif foreach:
raise RuntimeError(
f"foreach=True was passed, but can't use the foreach API on {device.type} tensors"
)
else:
clip_coef_clamped_device = clip_coef_clamped.to(device)
for g in device_grads:
g.mul_(clip_coef_clamped_device)
def _get_total_norm(
tensors: Union[torch.Tensor, List[torch.Tensor]],
norm_type: float = 2.0,
error_if_nonfinite: bool = False,
foreach: Optional[bool] = None,
) -> torch.Tensor:
if isinstance(tensors, torch.Tensor):
tensors = [tensors]
else:
tensors = list(tensors)
norm_type = float(norm_type)
if len(tensors) == 0:
return torch.tensor(0.0)
first_device = tensors[0].device
grouped_tensors: dict[tuple[torch.device, torch.dtype],
tuple[list[list[torch.Tensor]], list[int]]] = (
_group_tensors_by_device_and_dtype(
[tensors] # type: ignore[list-item]
)) # type: ignore[assignment]
norms: List[torch.Tensor] = []
for (device, _), ([device_tensors], _) in grouped_tensors.items():
local_tensors = [
t.to_local()
if isinstance(t, torch.distributed.tensor.DTensor) else t
for t in device_tensors
]
if (foreach is None and _has_foreach_support(local_tensors, device)
) or (foreach and _device_has_foreach_support(device)):
norms.extend(torch._foreach_norm(local_tensors, norm_type))
elif foreach:
raise RuntimeError(
f"foreach=True was passed, but can't use the foreach API on {device.type} tensors"
)
else:
norms.extend(
[torch.linalg.vector_norm(g, norm_type) for g in local_tensors])
total_norm = torch.linalg.vector_norm(
torch.stack([norm.to(first_device) for norm in norms]), norm_type)
if error_if_nonfinite and torch.logical_or(total_norm.isnan(),
total_norm.isinf()):
raise RuntimeError(
f"The total norm of order {norm_type} for gradients from "
"`parameters` is non-finite, so it cannot be clipped. To disable "
"this error and scale the gradients by the non-finite norm anyway, "
"set `error_if_nonfinite=False`")
return total_norm
def _get_foreach_kernels_supported_devices() -> list[str]:
r"""Return the device type list that supports foreach kernels."""
return ["cuda", "xpu", torch._C._get_privateuse1_backend_name()]
@torch.no_grad()
def _group_tensors_by_device_and_dtype(
tensorlistlist: List[List[Optional[torch.Tensor]]],
with_indices: bool = False,
) -> dict[tuple[torch.device, torch.dtype], tuple[
List[List[Optional[torch.Tensor]]], List[int]]]:
return torch._C._group_tensors_by_device_and_dtype(tensorlistlist,
with_indices)
def _device_has_foreach_support(device: torch.device) -> bool:
return device.type in (_get_foreach_kernels_supported_devices() +
["cpu"]) and not torch.jit.is_scripting()
def _has_foreach_support(tensors: List[torch.Tensor],
device: torch.device) -> bool:
return _device_has_foreach_support(device) and all(
t is None or type(t) in [torch.Tensor] for t in tensors)
@@ -1,19 +0,0 @@
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
class WanLatentPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
# def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
pass
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs):
logger.info("WAN Latent Pipeline forward")
pass

Some files were not shown because too many files have changed in this diff Show More