Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d0c53871a4 | ||
|
|
0d5306f61f |
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
@@ -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()),
|
||||
])
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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="")
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,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())
|
||||
|
||||
@@ -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()
|
||||
@@ -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
@@ -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
@@ -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,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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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),
|
||||
|
||||
@@ -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__()
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
"""
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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())):
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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.
|
||||
"""
|
||||
|
||||
@@ -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)
|
||||
@@ -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
Reference in New Issue
Block a user