Compare commits

..
119 changed files with 3239 additions and 938 deletions
+1 -1
View File
@@ -8,7 +8,7 @@ body:
attributes:
label: Environment
description: |
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
Please share your environment with us. You can run the command **python fastvideo/utils/env_utils.py** and copy-paste its output below.
placeholder: FastVideo version, platform, python version, cuda version...
validations:
required: true
+2 -2
View File
@@ -141,8 +141,8 @@ jobs:
fail-fast: false
matrix:
python-version: [
# {version: "3.10", tag: "latest"},
# {version: "3.11", tag: "py3.11-latest"},
{version: "3.10", tag: "latest"},
{version: "3.11", tag: "py3.11-latest"},
{version: "3.12", tag: "py3.12-latest"}
]
uses: ./.github/workflows/runpod-test.yml
+1
View File
@@ -20,6 +20,7 @@ exclude: |
fastvideo/train\.py|
fastvideo/utils/.*|
examples/.*|
fastvideo/v1/models/schedulers/scheduling_flow_match_euler_discrete.py|
.github/workflows/fastvideo-publish.yml|
.github/workflows/sta-publish.yml|
.github/workflows/build-image-template.yml|
@@ -70,8 +70,6 @@ DEFAULT_CONDA_PATTERNS = {
"optree",
"nccl",
"transformers",
"accelerate",
"peft",
"zmq",
"nvidia",
"pynvml",
@@ -87,8 +85,6 @@ DEFAULT_PIP_PATTERNS = {
"onnx",
"nccl",
"transformers",
"accelerate",
"peft",
"zmq",
"nvidia",
"pynvml",
+2 -1
View File
@@ -2,6 +2,7 @@ import torch
from flex_sta_ref import get_sliding_tile_attention_mask
from st_attn import sliding_tile_attention
from torch.nn.attention.flex_attention import flex_attention
# from flash_attn_interface import flash_attn_func
from tqdm import tqdm
flex_attention = torch.compile(flex_attention, dynamic=False)
@@ -22,7 +23,7 @@ def h100_fwd_kernel_test(Q, K, V, kernel_size):
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.linalg.norm(tensor, dim=-1, keepdim=True)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
+3 -1
View File
@@ -18,6 +18,7 @@ import os
import re
import sys
from pathlib import Path
from typing import Optional
import requests
@@ -167,7 +168,8 @@ _cached_base: str = ""
_cached_branch: str = ""
def get_repo_base_and_branch(pr_number: str) -> tuple[str | None, str | None]:
def get_repo_base_and_branch(
pr_number: str) -> tuple[Optional[str], Optional[str]]:
global _cached_base, _cached_branch
if _cached_base and _cached_branch:
return _cached_base, _cached_branch
+2 -1
View File
@@ -5,6 +5,7 @@ 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 = '../../../..'
@@ -88,7 +89,7 @@ class Example:
generate() -> str: Generates the documentation content.
""" # noqa: E501
path: Path
category: str | None = None
category: Optional[str] = None
main_file: Path = field(init=False)
other_files: list[Path] = field(init=False)
title: str = field(init=False)
+1 -2
View File
@@ -1,6 +1,5 @@
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
from fastvideo.version import __version__
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam"]
+2 -2
View File
@@ -237,7 +237,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
type=str,
default="540p",
choices=["540p", "720p"],
help="The resolution of the model.",
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--load-key",
@@ -361,7 +361,7 @@ def add_parallel_args(parser: argparse.ArgumentParser):
"--ring-degree",
type=int,
default=1,
help="Ring degree.",
help="Ulysses degree.",
)
return parser
+1 -1
View File
@@ -17,7 +17,7 @@ from fastvideo.models.hunyuan.vae import load_vae
from fastvideo.utils.parallel_states import nccl_info
class Inference:
class Inference(object):
def __init__(
self,
+1 -1
View File
@@ -41,7 +41,7 @@ def get_rewrite_prompt(ori_prompt, mode="Normal"):
elif mode == "Master":
prompt = master_mode_prompt.format(input=ori_prompt)
else:
raise Exception("Only supports Normal and Master mode, but got {}".format(mode))
raise Exception("Only supports Normal and Normal", mode)
return prompt
@@ -267,25 +267,25 @@ class Step1Model(PreTrainedModel):
class STEP1TextEncoder(torch.nn.Module):
def __init__(self, model_dir, max_length=320):
super()
super(STEP1TextEncoder, self).__init__()
self.max_length = max_length
self.text_tokenizer = Wrapped_StepChatTokenizer(os.path.join(model_dir, 'step1_chat_tokenizer.model'))
text_encoder = Step1Model.from_pretrained(model_dir)
self.text_encoder = text_encoder.eval().to(torch.bfloat16)
@torch.no_grad
@torch.autocast(device_type='cuda', dtype=torch.bfloat16)
def forward(self, prompts, with_mask=True, max_length=None):
self.device = next(self.text_encoder.parameters()).device
if type(prompts) is str:
prompts = [prompts]
with torch.no_grad(), torch.cuda.amp.autocast(dtype=torch.bfloat16):
if type(prompts) is str:
prompts = [prompts]
txt_tokens = self.text_tokenizer(prompts,
max_length=max_length or self.max_length,
padding="max_length",
truncation=True,
return_tensors="pt")
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
txt_tokens = self.text_tokenizer(prompts,
max_length=max_length or self.max_length,
padding="max_length",
truncation=True,
return_tensors="pt")
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
attention_mask=txt_tokens.attention_mask.to(self.device) if with_mask else None)
y_mask = txt_tokens.attention_mask
y_mask = txt_tokens.attention_mask
return y.transpose(0, 1), y_mask
+38
View File
@@ -0,0 +1,38 @@
import platform
import accelerate
import peft
import torch
import transformers
from transformers.utils import is_torch_cuda_available, is_torch_npu_available
VERSION = "1.2.0"
if __name__ == "__main__":
info = {
"FastVideo version": VERSION,
"Platform": platform.platform(),
"Python version": platform.python_version(),
"PyTorch version": torch.__version__,
"Transformers version": transformers.__version__,
"Accelerate version": accelerate.__version__,
"PEFT version": peft.__version__,
}
if is_torch_cuda_available():
info["PyTorch version"] += " (GPU)"
info["GPU type"] = torch.cuda.get_device_name()
if is_torch_npu_available():
info["PyTorch version"] += " (NPU)"
info["NPU type"] = torch.npu.get_device_name()
info["CANN version"] = torch.version.cann # codespell:ignore
try:
import bitsandbytes
info["Bitsandbytes version"] = bitsandbytes.__version__
except Exception:
pass
print("\n" + "\n".join([f"- {key}: {value}" for key, value in info.items()]) + "\n")
+8 -6
View File
@@ -3,7 +3,8 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass, fields
from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar
from typing import (TYPE_CHECKING, Any, Dict, Generic, Optional, Protocol, Set,
Type, TypeVar)
if TYPE_CHECKING:
from fastvideo.v1.fastvideo_args import FastVideoArgs
@@ -26,12 +27,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
@@ -45,7 +46,7 @@ class AttentionBackend(ABC):
@staticmethod
@abstractmethod
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
def get_builder_cls() -> Type["AttentionMetadataBuilder"]:
raise NotImplementedError
@@ -56,7 +57,8 @@ class AttentionMetadata:
current_timestep: int
def asdict_zerocopy(self,
skip_fields: set[str] | None = None) -> dict[str, Any]:
skip_fields: Optional[Set[str]] = None
) -> Dict[str, Any]:
"""Similar to dataclasses.asdict, but avoids deepcopying."""
if skip_fields is None:
skip_fields = set()
@@ -122,7 +124,7 @@ class AttentionImpl(ABC, Generic[T]):
head_size: int,
softmax_scale: float,
causal: bool = False,
num_kv_heads: int | None = None,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
@@ -1,5 +1,7 @@
# 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
@@ -26,7 +28,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
@@ -34,15 +36,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
@@ -54,7 +56,7 @@ class FlashAttentionImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
+5 -3
View File
@@ -1,3 +1,5 @@
from typing import List, Optional, Type
import torch
from sageattention import sageattn
@@ -15,7 +17,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
@@ -23,7 +25,7 @@ class SageAttentionBackend(AttentionBackend):
return "SAGE_ATTN"
@staticmethod
def get_impl_cls() -> type["SageAttentionImpl"]:
def get_impl_cls() -> Type["SageAttentionImpl"]:
return SageAttentionImpl
# @staticmethod
@@ -39,7 +41,7 @@ class SageAttentionImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
+5 -3
View File
@@ -1,3 +1,5 @@
from typing import List, Optional, Type
import torch
from fastvideo.v1.attention.backends.abstract import (
@@ -14,7 +16,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
@@ -22,7 +24,7 @@ class SDPABackend(AttentionBackend):
return "SDPA"
@staticmethod
def get_impl_cls() -> type["SDPAImpl"]:
def get_impl_cls() -> Type["SDPAImpl"]:
return SDPAImpl
# @staticmethod
@@ -38,7 +40,7 @@ class SDPAImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
@@ -1,5 +1,6 @@
import json
from dataclasses import dataclass
from typing import List, Optional, Type
import torch
from einops import rearrange
@@ -19,7 +20,7 @@ logger = init_logger(__name__)
# TODO(will-refactor): move this to a utils file
def dict_to_3d_list(mask_strategy) -> list[list[list[torch.Tensor | None]]]:
def dict_to_3d_list(mask_strategy) -> List[List[List[Optional[torch.Tensor]]]]:
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
max_timesteps_idx = max(
@@ -57,7 +58,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]
@@ -66,15 +67,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
@@ -109,7 +110,7 @@ class SlidingTileAttentionImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
+14 -12
View File
@@ -1,5 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional, Tuple
import torch
import torch.nn as nn
@@ -20,11 +22,11 @@ class DistributedAttention(nn.Module):
def __init__(self,
num_heads: int,
head_size: int,
num_kv_heads: int | None = None,
softmax_scale: float | None = None,
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: tuple[_Backend, ...]
| None = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = "",
**extra_impl_args) -> None:
super().__init__()
@@ -60,10 +62,10 @@ class DistributedAttention(nn.Module):
q: torch.Tensor,
k: torch.Tensor,
v: 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]:
replicated_q: Optional[torch.Tensor] = None,
replicated_k: Optional[torch.Tensor] = None,
replicated_v: Optional[torch.Tensor] = None,
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
"""Forward pass for distributed attention.
Args:
@@ -139,11 +141,11 @@ class LocalAttention(nn.Module):
def __init__(self,
num_heads: int,
head_size: int,
num_kv_heads: int | None = None,
softmax_scale: float | None = None,
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: tuple[_Backend, ...]
| None = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
**extra_impl_args) -> None:
super().__init__()
if softmax_scale is None:
+13 -14
View File
@@ -2,10 +2,9 @@
# 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 cast
from typing import Generator, Optional, Tuple, Type, cast
import torch
@@ -18,7 +17,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) -> _Backend | None:
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
"""
Convert a string backend name to a _Backend enum value.
@@ -32,7 +31,7 @@ def backend_name_to_enum(backend_name: str) -> _Backend | None:
None
def get_env_variable_attn_backend() -> _Backend | None:
def get_env_variable_attn_backend() -> Optional[_Backend]:
'''
Get the backend override specified by the FastVideo attention
backend environment variable, if one is specified.
@@ -54,10 +53,10 @@ def get_env_variable_attn_backend() -> _Backend | None:
#
# THIS SELECTION TAKES PRECEDENCE OVER THE
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
forced_attn_backend: _Backend | None = None
forced_attn_backend: Optional[_Backend] = None
def global_force_attn_backend(attn_backend: _Backend | None) -> None:
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
'''
Force all attention operations to use a specified backend.
@@ -72,7 +71,7 @@ def global_force_attn_backend(attn_backend: _Backend | None) -> None:
forced_attn_backend = attn_backend
def get_global_forced_attn_backend() -> _Backend | None:
def get_global_forced_attn_backend() -> Optional[_Backend]:
'''
Get the currently-forced choice of attention backend,
or None if auto-selection is currently enabled.
@@ -83,8 +82,8 @@ def get_global_forced_attn_backend() -> _Backend | None:
def get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: tuple[_Backend, ...] | None = None,
) -> type[AttentionBackend]:
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
) -> Type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype,
supported_attention_backends)
@@ -93,8 +92,8 @@ def get_attn_backend(
def _cached_get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: tuple[_Backend, ...] | None = None,
) -> type[AttentionBackend]:
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
) -> Type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
#
@@ -103,13 +102,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: _Backend | None = (
backend_by_global_setting: Optional[_Backend] = (
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: str | None = envs.FASTVIDEO_ATTENTION_BACKEND
backend_by_env_var: Optional[str] = envs.FASTVIDEO_ATTENTION_BACKEND
if backend_by_env_var is not None:
selected_backend = backend_name_to_enum(backend_by_env_var)
@@ -121,7 +120,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
+11 -3
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field, fields
from typing import Any
from typing import Any, Dict
from fastvideo.v1.logger import init_logger
@@ -41,7 +41,15 @@ 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:
# Remove all keys that start with "_"
keys_to_remove = [
key for key in list(source_model_dict.keys())
if str(key).startswith("_")
]
for key in keys_to_remove:
source_model_dict.pop(key)
arch_config = self.arch_config
valid_fields = {f.name for f in fields(arch_config)}
@@ -55,7 +63,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)}
+4 -1
View File
@@ -1,5 +1,8 @@
from fastvideo.v1.configs.models.dits.flux import FluxImageConfig
from fastvideo.v1.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.v1.configs.models.dits.stepvideo import StepVideoConfig
from fastvideo.v1.configs.models.dits.wanvideo import WanVideoConfig
__all__ = ["HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig"]
__all__ = [
"HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig", "FluxImageConfig"
]
+3 -3
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Any
from typing import Any, Optional, Tuple
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: QuantizationConfig | None = None
quant_config: Optional[QuantizationConfig] = None
@staticmethod
def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any:
+130
View File
@@ -0,0 +1,130 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
import torch
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class FluxImageArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
_param_names_mapping: dict = field(
default_factory=lambda: {
# 1. context_embedder to txt_in mapping:
r"^context_embedder\.(.*)$":
r"txt_in.\1",
# 2. x_embedder to img_in mapping:
r"^x_embedder\.(.*)$":
r"img_in.\1",
# 3. Top-level time_text_embed mappings:
r"^time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"time_in.mlp.fc_in.\1",
r"^time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"time_in.mlp.fc_out.\1",
r"^time_text_embed\.guidance_embedder\.linear_1\.(.*)$":
r"guidance_in.mlp.fc_in.\1",
r"^time_text_embed\.guidance_embedder\.linear_2\.(.*)$":
r"guidance_in.mlp.fc_out.\1",
r"^time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"txt2_in.fc_in.\1",
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"txt2_in.fc_out.\1",
# 4. transformer_blocks mapping:
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
r"double_blocks.\1.img_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
r"double_blocks.\1.txt_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"double_blocks.\1.img_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"double_blocks.\1.img_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"double_blocks.\1.img_attn_proj.\2",
# Corrected: merge attn.to_add_out into the main projection.
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
r"double_blocks.\1.txt_attn_proj.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
r"double_blocks.\1.txt_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
r"double_blocks.\1.txt_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_out.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_out.\2",
# 5. single_transformer_blocks mapping:
r"^single_transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"single_blocks.\1.q_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"single_blocks.\1.k_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"single_blocks.\1.linear1.\2", 0, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"single_blocks.\1.linear1.\2", 1, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"single_blocks.\1.linear1.\2", 2, 4),
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$":
(r"single_blocks.\1.linear1.\2", 3, 4),
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$":
r"single_blocks.\1.linear2.\2",
r"^single_transformer_blocks\.(\d+)\.norm\.linear\.(.*)$":
r"single_blocks.\1.modulation.linear.\2",
# 6. Final layers mapping:
r"^norm_out\.linear\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
r"^proj_out\.(.*)$":
r"final_layer.linear.\1",
})
patch_size: int = 1
in_channels: int = 64
out_channels: Optional[int] = None
num_layers: int = 19
num_single_layers: int = 38
attention_head_dim: int = 128
num_attention_heads: int = 24
joint_attention_dim: int = 4096
pooled_projection_dim: int = 768
guidance_embeds: bool = False
axes_dims_rope: Tuple[int, ...] = (16, 56, 56)
rope_theta: int = 10000
dtype: Optional[torch.dtype] = torch.bfloat16
def __post_init__(self):
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.in_channels // 4
@dataclass
class FluxImageConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=FluxImageArchConfig)
prefix: str = "Flux"
@@ -1,4 +1,5 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
import torch
@@ -155,9 +156,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: torch.dtype | None = None
dtype: Optional[torch.dtype] = None
text_embed_dim: int = 4096
pooled_projection_dim: int = 768
rope_theta: int = 256
@@ -1,4 +1,5 @@
from dataclasses import dataclass, field
from typing import List, Optional, Tuple, Union
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -39,17 +40,17 @@ class StepVideoArchConfig(DiTArchConfig):
num_attention_heads: int = 48
attention_head_dim: int = 128
in_channels: int = 64
out_channels: int | None = 64
out_channels: Optional[int] = 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: int | list[int] | tuple[int, ...] | None = field(
caption_channels: Optional[Union[int, List[int], Tuple[int, ...]]] = field(
default_factory=lambda: [6144, 1024])
attention_type: str | None = "torch"
use_additional_conditions: bool | None = False
attention_type: Optional[str] = "torch"
use_additional_conditions: Optional[bool] = False
def __post_init__(self):
self.hidden_size = self.num_attention_heads * self.attention_head_dim
+4 -3
View File
@@ -1,4 +1,5 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -51,7 +52,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
@@ -64,8 +65,8 @@ class WanVideoArchConfig(DiTArchConfig):
cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6
image_dim: int | None = None
added_kv_proj_dim: int | None = None
image_dim: Optional[int] = None
added_kv_proj_dim: Optional[int] = None
rope_max_seq_len: int = 1024
def __post_init__(self):
+11 -11
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Any
from typing import Any, Dict, List, Optional, Tuple
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: 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
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
@dataclass
@@ -61,8 +61,8 @@ class EncoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=EncoderArchConfig)
prefix: str = ""
quant_config: QuantizationConfig | None = None
lora_config: Any | None = None
quant_config: Optional[QuantizationConfig] = None
lora_config: Optional[Any] = None
@dataclass
+5 -4
View File
@@ -1,4 +1,5 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
ImageEncoderConfig,
@@ -50,8 +51,8 @@ class CLIPTextConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(
default_factory=CLIPTextArchConfig)
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
num_hidden_layers_override: Optional[int] = None
require_post_norm: Optional[bool] = None
prefix: str = "clip"
@@ -60,6 +61,6 @@ class CLIPVisionConfig(ImageEncoderConfig):
arch_config: ImageEncoderArchConfig = field(
default_factory=CLIPVisionArchConfig)
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
num_hidden_layers_override: Optional[int] = None
require_post_norm: Optional[bool] = None
prefix: str = "clip"
@@ -1,4 +1,5 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
@@ -11,7 +12,7 @@ class LlamaArchConfig(TextEncoderArchConfig):
intermediate_size: int = 11008
num_hidden_layers: int = 32
num_attention_heads: int = 32
num_key_value_heads: int | None = None
num_key_value_heads: Optional[int] = None
hidden_act: str = "silu"
max_position_embeddings: int = 2048
initializer_range: float = 0.02
@@ -23,11 +24,11 @@ class LlamaArchConfig(TextEncoderArchConfig):
pretraining_tp: int = 1
tie_word_embeddings: bool = False
rope_theta: float = 10000.0
rope_scaling: float | None = None
rope_scaling: Optional[float] = None
attention_bias: bool = False
attention_dropout: float = 0.0
mlp_bias: bool = False
head_dim: int | None = None
head_dim: Optional[int] = None
hidden_state_skip_layer: int = 2
text_len: int = 256
+2 -1
View File
@@ -1,4 +1,5 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
@@ -11,7 +12,7 @@ class T5ArchConfig(TextEncoderArchConfig):
d_kv: int = 64
d_ff: int = 2048
num_layers: int = 6
num_decoder_layers: int | None = None
num_decoder_layers: Optional[int] = None
num_heads: int = 8
relative_attention_num_buckets: int = 32
relative_attention_max_distance: int = 128
@@ -1,4 +1,5 @@
from fastvideo.v1.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.v1.configs.models.vaes.image_vae import ImageVAEConfig
from fastvideo.v1.configs.models.vaes.stepvideovae import StepVideoVAEConfig
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
@@ -6,4 +7,5 @@ __all__ = [
"HunyuanVAEConfig",
"WanVAEConfig",
"StepVideoVAEConfig",
"ImageVAEConfig",
]
+2 -2
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Any
from typing import Any, Union
import torch
@@ -9,7 +9,7 @@ from fastvideo.v1.utils import StoreBoolean
@dataclass
class VAEArchConfig(ArchConfig):
scaling_factor: float | torch.Tensor = 0
scaling_factor: Union[float, torch.tensor] = 0
temporal_compression_ratio: int = 4
spatial_compression_ratio: int = 8
@@ -1,4 +1,5 @@
from dataclasses import dataclass, field
from typing import Tuple
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
@@ -8,19 +9,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
@@ -0,0 +1,41 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class ImageVAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
down_block_types: Tuple[str] = ("DownEncoderBlock2D", )
up_block_types: Tuple[str] = ("UpDecoderBlock2D", )
block_out_channels: Tuple[int] = (64, )
layers_per_block: int = 1
act_fn: str = "silu"
latent_channels: int = 4
norm_num_groups: int = 32
sample_size: int = 32
scaling_factor: float = 0.18215
shift_factor: Optional[float] = None
latents_mean: Optional[Tuple[float]] = None
latents_std: Optional[Tuple[float]] = None
force_upcast: float = True
use_quant_conv: bool = True
use_post_quant_conv: bool = True
mid_block_add_attention: bool = True
def __post_init__(self):
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels) -
1)
self.temporal_compression_ratio = 1
@dataclass
class ImageVAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=ImageVAEArchConfig)
# overrides VAEConfig
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
+8 -7
View File
@@ -1,4 +1,5 @@
from dataclasses import dataclass, field
from typing import Tuple
import torch
@@ -9,12 +10,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,
@@ -32,7 +33,7 @@ class WanVAEArchConfig(VAEArchConfig):
0.2503,
-0.2921,
)
latents_std: tuple[float, ...] = (
latents_std: Tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
@@ -54,9 +55,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)
+11 -13
View File
@@ -1,7 +1,6 @@
import json
from collections.abc import Callable
from dataclasses import asdict, dataclass, field, fields
from typing import Any, cast
from typing import Any, Callable, Dict, Optional, Tuple, cast
import torch
@@ -18,7 +17,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
@@ -27,7 +26,7 @@ class PipelineConfig:
"""Base configuration for all pipeline architectures."""
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: float | None = None
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
@@ -44,18 +43,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: str | None = None
mask_strategy_file_path: Optional[str] = None
# Compilation
enable_torch_compile: bool = False
@@ -108,7 +107,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
@@ -124,9 +123,8 @@ 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,
strict=False):
for target_config, source_config in zip(
current_value, new_value):
target_config.update_model_config(source_config)
else:
setattr(self, key, new_value)
+68
View File
@@ -0,0 +1,68 @@
from dataclasses import dataclass, field
from typing import Callable, Tuple
import torch
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.configs.models.dits import FluxImageConfig
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
CLIPTextConfig, T5Config)
from fastvideo.v1.configs.models.vaes import ImageVAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
def t5_preprocess_text(prompt: str) -> str:
return prompt
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
hidden_state: torch.tensor = outputs.last_hidden_state
assert torch.isnan(hidden_state).sum() == 0
prompt_embeds_tensor: torch.tensor = torch.stack([
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
for u in hidden_state
],
dim=0)
return prompt_embeds_tensor
def clip_preprocess_text(prompt: str) -> str:
return prompt
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
pooler_output: torch.tensor = outputs.pooler_output
return pooler_output
@dataclass
class FluxConfig(PipelineConfig):
"""Base configuration for Flux pipeline architecture."""
# FluxConfig-specific parameters with defaults
# DiT
dit_config: DiTConfig = field(default_factory=FluxImageConfig)
# VAE
vae_config: VAEConfig = field(default_factory=ImageVAEConfig)
# Denoising stage
embedded_cfg_scale: float = 3.5
# Text encoding stage
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (CLIPTextConfig(), T5Config()))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (clip_preprocess_text, t5_preprocess_text))
postprocess_text_funcs: Tuple[
Callable[[BaseEncoderOutput], torch.tensor],
...] = field(default_factory=lambda:
(clip_postprocess_text, t5_postprocess_text))
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("bf16", "bf16"))
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
+11 -12
View File
@@ -1,6 +1,5 @@
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import TypedDict
from typing import Callable, Tuple, TypedDict
import torch
@@ -36,11 +35,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:]
@@ -51,8 +50,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
@@ -73,19 +72,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):
+11 -6
View File
@@ -1,9 +1,10 @@
"""Registry for pipeline weight-specific configurations."""
import os
from collections.abc import Callable
from typing import Callable, Dict, Optional, Type
from fastvideo.v1.configs.pipelines.base import PipelineConfig
from fastvideo.v1.configs.pipelines.flux import FluxConfig
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
HunyuanConfig)
from fastvideo.v1.configs.pipelines.stepvideo import StepVideoT2VConfig
@@ -18,7 +19,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,
@@ -26,32 +27,36 @@ WEIGHT_CONFIG_REGISTRY: dict[str, type[PipelineConfig]] = {
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V720PConfig,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
"black-forest-labs/FLUX.1-dev": FluxConfig,
"black-forest-labs/FLUX.1-schnell": FluxConfig,
# Add other specific weight variants
}
# 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(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
"flux": lambda id: "flux" in id.lower(),
# Add other pipeline architecture detectors
}
# 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":
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
"stepvideo": StepVideoT2VConfig
"stepvideo": StepVideoT2VConfig,
"flux": FluxConfig,
# Other fallbacks by architecture
}
def get_pipeline_config_cls_for_name(
pipeline_name_or_path: str) -> type[PipelineConfig] | None:
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
"""Get the appropriate config class for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
+9 -11
View File
@@ -1,5 +1,5 @@
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Callable, Tuple
import torch
@@ -11,15 +11,13 @@ 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, strict=False)
]
prompt_embeds_tensor: torch.Tensor = torch.stack([
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens)]
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
],
@@ -46,16 +44,16 @@ class WanT2V480PConfig(PipelineConfig):
flow_shift: int = 3
# Text encoding stage
text_encoder_configs: tuple[EncoderConfig, ...] = field(
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (T5Config(), ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
...] = field(default_factory=lambda:
(t5_postprocess_text, ))
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: tuple[str, ...] = field(
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp32", ))
# WanConfig-specific added parameters
+6 -6
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass
from typing import Any
from typing import Any, Dict, List, Optional, Union
from fastvideo.v1.logger import init_logger
@@ -15,12 +15,12 @@ class SamplingParam:
data_type: str = "video"
# Image inputs
image_path: str | None = None
image_path: Optional[str] = None
# Text inputs
prompt: str | list[str] | None = None
negative_prompt: str | None = None
prompt_path: str | None = None
prompt: Optional[Union[str, List[str]]] = None
negative_prompt: Optional[str] = None
prompt_path: Optional[str] = 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)
+14
View File
@@ -0,0 +1,14 @@
from dataclasses import dataclass
from fastvideo.v1.configs.sample.base import SamplingParam
@dataclass
class FluxSamplingParam(SamplingParam):
# Video parameters
height: int = 1024
width: int = 1024
num_frames: int = 1
# Denoising stage
num_inference_steps: int = 50
+12 -7
View File
@@ -1,7 +1,7 @@
import os
from collections.abc import Callable
from typing import Any
from typing import Any, Callable, Dict, Optional
from fastvideo.v1.configs.sample.flux import FluxSamplingParam
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.v1.configs.sample.stepvideo import StepVideoT2VSamplingParam
@@ -15,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,
@@ -23,31 +23,36 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
"black-forest-labs/FLUX.1-dev": FluxSamplingParam,
"black-forest-labs/FLUX.1-schnell": FluxSamplingParam,
# Add other specific weight variants
}
# 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(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
"flux": lambda id: "flux" in id.lower(),
# Add other pipeline architecture detectors
}
# 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":
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
"stepvideo": StepVideoT2VSamplingParam
"stepvideo": StepVideoT2VSamplingParam,
"flux": FluxSamplingParam,
# Other fallbacks by architecture
}
def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
def get_sampling_param_cls_for_name(
pipeline_name_or_path: str) -> Optional[Any]:
"""Get the appropriate sampling param for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
@@ -1,6 +1,8 @@
# 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 Optional
import torch
import torch.distributed as dist
from torch.distributed import ProcessGroup
@@ -16,8 +18,8 @@ class DeviceCommunicatorBase:
def __init__(self,
cpu_group: ProcessGroup,
device: torch.device | None = None,
device_group: ProcessGroup | None = None,
device: Optional[torch.device] = None,
device_group: Optional[ProcessGroup] = None,
unique_name: str = ""):
self.device = device or torch.device("cpu")
self.cpu_group = cpu_group
@@ -64,7 +66,7 @@ class DeviceCommunicatorBase:
def gather(self,
input_: torch.Tensor,
dst: int = 0,
dim: int = -1) -> torch.Tensor | None:
dim: int = -1) -> Optional[torch.Tensor]:
"""
NOTE: We assume that the input tensor is on the same device across
all the ranks.
@@ -168,7 +170,7 @@ class DeviceCommunicatorBase:
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:
def send(self, tensor: torch.Tensor, dst: Optional[int] = 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:
@@ -178,7 +180,7 @@ class DeviceCommunicatorBase:
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: int | None = None) -> torch.Tensor:
src: Optional[int] = 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,6 +1,8 @@
# 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
@@ -12,15 +14,15 @@ class CudaCommunicator(DeviceCommunicatorBase):
def __init__(self,
cpu_group: ProcessGroup,
device: torch.device | None = None,
device_group: ProcessGroup | None = None,
device: Optional[torch.device] = None,
device_group: Optional[ProcessGroup] = 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: PyNcclCommunicator | None = None
self.pynccl_comm: Optional[PyNcclCommunicator] = None
if self.world_size > 1:
self.pynccl_comm = PyNcclCommunicator(
group=self.cpu_group,
@@ -40,7 +42,7 @@ class CudaCommunicator(DeviceCommunicatorBase):
torch.distributed.all_reduce(out, group=self.device_group)
return out
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
def send(self, tensor: torch.Tensor, dst: Optional[int] = 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:
@@ -55,7 +57,7 @@ class CudaCommunicator(DeviceCommunicatorBase):
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: int | None = None) -> torch.Tensor:
src: Optional[int] = 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,6 +1,8 @@
# 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
@@ -20,9 +22,9 @@ class PyNcclCommunicator:
def __init__(
self,
group: ProcessGroup | StatelessProcessGroup,
device: int | str | torch.device,
library_path: str | None = None,
group: Union[ProcessGroup, StatelessProcessGroup],
device: Union[int, str, torch.device],
library_path: Optional[str] = None,
):
"""
Args:
@@ -27,7 +27,7 @@
import ctypes
import platform
from dataclasses import dataclass
from typing import Any
from typing import Any, Dict, List, Optional
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: str | None = None):
def __init__(self, so_file: Optional[str] = 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
+44 -45
View File
@@ -27,11 +27,10 @@ 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, Optional
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
from unittest.mock import patch
import torch
@@ -58,15 +57,15 @@ TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
def _split_tensor_dict(
tensor_dict: dict[str, torch.Tensor | Any]
) -> tuple[list[tuple[str, Any]], list[torch.Tensor]]:
tensor_dict: Dict[str, Union[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,
@@ -82,7 +81,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:
@@ -98,7 +97,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:
@@ -129,7 +128,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:
@@ -144,16 +143,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: Any | None # shared memory broadcaster
mq_broadcaster: Optional[Any] # shared memory broadcaster
def __init__(
self,
group_ranks: list[list[int]],
group_ranks: List[List[int]],
local_rank: int,
torch_distributed_backend: str | Backend,
torch_distributed_backend: Union[str, Backend],
use_device_communicator: bool,
use_message_queue_broadcaster: bool = False,
group_name: str | None = None,
group_name: Optional[str] = None,
):
group_name = group_name or "anonymous"
self.unique_name = _get_unique_name(group_name)
@@ -244,8 +243,8 @@ class GroupCoordinator:
return self.ranks[(rank_in_group - 1) % world_size]
@contextmanager
def graph_capture(self,
graph_capture_context: GraphCaptureContext | None = None):
def graph_capture(
self, graph_capture_context: Optional[GraphCaptureContext] = None):
if graph_capture_context is None:
stream = torch.cuda.Stream()
graph_capture_context = GraphCaptureContext(stream)
@@ -302,7 +301,7 @@ class GroupCoordinator:
def gather(self,
input_: torch.Tensor,
dst: int = 0,
dim: int = -1) -> torch.Tensor | None:
dim: int = -1) -> Optional[torch.Tensor]:
"""
NOTE: We assume that the input tensor is on the same device across
all the ranks.
@@ -338,7 +337,7 @@ class GroupCoordinator:
group=self.device_group)
return input_
def broadcast_object(self, obj: Any | None = None, src: int = 0):
def broadcast_object(self, obj: Optional[Any] = None, src: int = 0):
"""Broadcast the input object.
NOTE: `src` is the local rank of the source rank.
"""
@@ -363,9 +362,9 @@ class GroupCoordinator:
return recv[0]
def broadcast_object_list(self,
obj_list: list[Any],
obj_list: List[Any],
src: int = 0,
group: ProcessGroup | None = None):
group: Optional[ProcessGroup] = None):
"""Broadcast the input object list.
NOTE: `src` is the local rank of the source rank.
"""
@@ -445,11 +444,11 @@ class GroupCoordinator:
def broadcast_tensor_dict(
self,
tensor_dict: dict[str, torch.Tensor | Any] | None = None,
tensor_dict: Optional[Dict[str, Union[torch.Tensor, Any]]] = None,
src: int = 0,
group: ProcessGroup | None = None,
metadata_group: ProcessGroup | None = None
) -> dict[str, torch.Tensor | Any] | None:
group: Optional[ProcessGroup] = None,
metadata_group: Optional[ProcessGroup] = None
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
"""Broadcast the input tensor dictionary.
NOTE: `src` is the local rank of the source rank.
"""
@@ -463,7 +462,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)}")
@@ -530,10 +529,10 @@ class GroupCoordinator:
def send_tensor_dict(
self,
tensor_dict: dict[str, torch.Tensor | Any],
dst: int | None = None,
tensor_dict: Dict[str, Union[torch.Tensor, Any]],
dst: Optional[int] = None,
all_gather_group: Optional["GroupCoordinator"] = None,
) -> dict[str, torch.Tensor | Any] | None:
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
"""Send the input tensor dictionary.
NOTE: `dst` is the local rank of the source rank.
"""
@@ -553,7 +552,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)}"
@@ -584,9 +583,9 @@ class GroupCoordinator:
def recv_tensor_dict(
self,
src: int | None = None,
src: Optional[int] = None,
all_gather_group: Optional["GroupCoordinator"] = None,
) -> dict[str, torch.Tensor | Any] | None:
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
"""Recv the input tensor dictionary.
NOTE: `src` is the local rank of the source rank.
"""
@@ -607,7 +606,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,
@@ -657,7 +656,7 @@ class GroupCoordinator:
"""
torch.distributed.barrier(group=self.cpu_group)
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
def send(self, tensor: torch.Tensor, dst: Optional[int] = 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)
@@ -665,7 +664,7 @@ class GroupCoordinator:
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: int | None = None) -> torch.Tensor:
src: Optional[int] = 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)
@@ -683,7 +682,7 @@ class GroupCoordinator:
self.mq_broadcaster = None
_WORLD: GroupCoordinator | None = None
_WORLD: Optional[GroupCoordinator] = None
def get_world_group() -> GroupCoordinator:
@@ -691,7 +690,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],
@@ -703,11 +702,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: str | None = None,
group_name: Optional[str] = None,
) -> GroupCoordinator:
return GroupCoordinator(
@@ -720,7 +719,7 @@ def init_model_parallel_group(
)
_TP: GroupCoordinator | None = None
_TP: Optional[GroupCoordinator] = None
def get_tp_group() -> GroupCoordinator:
@@ -779,7 +778,7 @@ def init_distributed_environment(
"world group already initialized with a different world size")
_SP: GroupCoordinator | None = None
_SP: Optional[GroupCoordinator] = None
def get_sp_group() -> GroupCoordinator:
@@ -790,7 +789,7 @@ def get_sp_group() -> GroupCoordinator:
def initialize_model_parallel(
tensor_model_parallel_size: int = 1,
sequence_model_parallel_size: int = 1,
backend: str | None = None,
backend: Optional[str] = None,
) -> None:
"""
Initialize model parallel groups.
@@ -859,7 +858,7 @@ def get_sequence_model_parallel_rank() -> int:
def ensure_model_parallel_initialized(
tensor_model_parallel_size: int,
sequence_model_parallel_size: int,
backend: str | None = None,
backend: Optional[str] = None,
) -> None:
"""Helper to initialize model parallel groups if they are not initialized,
or ensure tensor-parallel, sequence-parallel sizes
@@ -970,8 +969,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: ProcessGroup | StatelessProcessGroup,
source_rank: int = 0) -> list[bool]:
def in_the_same_node_as(pg: Union[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
@@ -1057,7 +1056,7 @@ def in_the_same_node_as(pg: ProcessGroup | StatelessProcessGroup,
def initialize_tensor_parallel_group(
tensor_model_parallel_size: int = 1,
backend: str | None = None,
backend: Optional[str] = None,
group_name_suffix: str = "") -> GroupCoordinator:
"""Initialize a tensor parallel group for a specific model.
@@ -1121,7 +1120,7 @@ def initialize_tensor_parallel_group(
def initialize_sequence_parallel_group(
sequence_model_parallel_size: int = 1,
backend: str | None = None,
backend: Optional[str] = None,
group_name_suffix: str = "") -> GroupCoordinator:
"""Initialize a sequence parallel group for a specific model.
+6 -7
View File
@@ -9,8 +9,7 @@ import dataclasses
import pickle
import time
from collections import deque
from collections.abc import Sequence
from typing import Any
from typing import Any, Deque, Dict, Optional, Sequence, Tuple
import torch
from torch.distributed import TCPStore
@@ -73,15 +72,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
@@ -115,7 +114,7 @@ class StatelessProcessGroup:
self.recv_src_counter[src] += 1
return obj
def broadcast_obj(self, obj: Any | None, src: int) -> Any:
def broadcast_obj(self, obj: Optional[Any], src: int) -> Any:
"""Broadcast an object from a source rank to all other ranks.
It does not clean up after all ranks have received the object.
Use it for limited times, e.g., for initialization.
+6 -6
View File
@@ -4,7 +4,7 @@
import argparse
import dataclasses
import os
from typing import Any, cast
from typing import Any, Dict, List, Optional, 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: str | None = None) -> None:
args_dict: Dict[str, Any],
prefix: Optional[str] = None) -> None:
"""
Update configuration object from arguments dictionary.
+3 -1
View File
@@ -1,12 +1,14 @@
# 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())
+3 -2
View File
@@ -4,6 +4,7 @@ import argparse
import os
import subprocess
import sys
from typing import List, Optional
from fastvideo.v1.logger import init_logger
@@ -18,8 +19,8 @@ class RaiseNotImplementedAction(argparse.Action):
def launch_distributed(num_gpus: int,
args: list[str],
master_port: int | None = None) -> int:
args: List[str],
master_port: Optional[int] = None) -> int:
"""
Launch a distributed job with the given arguments
+133 -7
View File
@@ -10,7 +10,7 @@ import gc
import math
import os
import time
from typing import Any
from typing import Any, Dict, List, Optional, Union
import imageio
import numpy as np
@@ -53,9 +53,11 @@ class VideoGenerator:
@classmethod
def from_pretrained(cls,
model_path: str,
device: str | None = None,
torch_dtype: torch.dtype | None = None,
pipeline_config: str | PipelineConfig | None = None,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[
Union[str
| PipelineConfig]] = None,
**kwargs) -> "VideoGenerator":
"""
Create a video generator from a pretrained model.
@@ -126,9 +128,9 @@ class VideoGenerator:
def generate_video(
self,
prompt: str,
sampling_param: SamplingParam | None = None,
sampling_param: Optional[SamplingParam] = None,
**kwargs,
) -> dict[str, Any] | list[np.ndarray]:
) -> Union[Dict[str, Any], List[np.ndarray]]:
"""
Generate a video based on the given prompt.
@@ -259,7 +261,7 @@ class VideoGenerator:
# Run inference
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch
samples = output_batch.output.cpu()
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
@@ -293,6 +295,130 @@ class VideoGenerator:
"generation_time": gen_time
}
def generate_image(
self,
prompt: str,
sampling_param: Optional[SamplingParam] = None,
**kwargs,
) -> Union[Dict[str, Any], List[np.ndarray]]:
"""
Generate a image based on the given prompt.
Args:
prompt: The prompt to use for generation
negative_prompt: The negative prompt to use (overrides the one in fastvideo_args)
output_path: Path to save the image (overrides the one in fastvideo_args)
save_video: Whether to save the image to disk
return_frames: Whether to return the raw frames
num_inference_steps: Number of denoising steps (overrides fastvideo_args)
guidance_scale: Classifier-free guidance scale (overrides fastvideo_args)
height: Height of generated image (overrides fastvideo_args)
width: Width of generated image (overrides fastvideo_args)
seed: Random seed for generation (overrides fastvideo_args)
callback: Callback function called after each step
callback_steps: Number of steps between each callback
Returns:
Either the output dictionary or the list of frames depending on return_frames
"""
# Create a copy of inference args to avoid modifying the original
fastvideo_args = self.fastvideo_args
# Validate inputs
if not isinstance(prompt, str):
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
kwargs["prompt"] = prompt
sampling_param.update(kwargs)
# Process negative prompt
if sampling_param.negative_prompt is not None:
sampling_param.negative_prompt = sampling_param.negative_prompt.strip(
)
# Validate dimensions
if (sampling_param.height <= 0 or sampling_param.width <= 0
or sampling_param.num_frames != 1):
raise ValueError(
f"Height, width must be positive integers, num_frames must be 1, got "
f"height={sampling_param.height}, width={sampling_param.width}, "
f"num_frames={sampling_param.num_frames}")
# Calculate sizes
target_height = align_to(sampling_param.height, 16)
target_width = align_to(sampling_param.width, 16)
# Calculate latent sizes
latents_size = [
1, sampling_param.height // 8, sampling_param.width // 8
]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# Log parameters
debug_str = f"""
height: {target_height}
width: {target_width}
prompt: {prompt}
neg_prompt: {sampling_param.negative_prompt}
seed: {sampling_param.seed}
infer_steps: {sampling_param.num_inference_steps}
num_images_per_prompt: {sampling_param.num_videos_per_prompt}
guidance_scale: {sampling_param.guidance_scale}
n_tokens: {n_tokens}
flow_shift: {fastvideo_args.flow_shift}
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}
save: {sampling_param.save_video}
output_path: {sampling_param.output_path}
""" # type: ignore[attr-defined]
logger.info(debug_str)
# Prepare batch
batch = ForwardBatch(
**shallow_asdict(sampling_param),
eta=0.0,
n_tokens=n_tokens,
extra={},
)
# Run inference
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch.output.cpu()
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
# Process outputs
x = torchvision.utils.make_grid(samples, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
images = (x * 255).numpy().astype(np.uint8)
# Save video if requested
if batch.save_video:
save_path = batch.output_path
if save_path:
os.makedirs(os.path.dirname(save_path), exist_ok=True)
image_path = os.path.join(save_path, f"{prompt[:100]}.png")
imageio.imwrite(image_path, images)
logger.info("Saved image to %s", image_path)
else:
logger.warning("No output path provided, image not saved")
if batch.return_frames:
return [images]
else:
return {
"samples": samples,
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time
}
def shutdown(self):
"""
Shutdown the video generator.
+12 -13
View File
@@ -2,29 +2,28 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/envs.py
import os
from collections.abc import Callable
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, Callable, Dict, Optional
if TYPE_CHECKING:
FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60
FASTVIDEO_NCCL_SO_PATH: str | None = None
LD_LIBRARY_PATH: str | None = None
FASTVIDEO_NCCL_SO_PATH: Optional[str] = None
LD_LIBRARY_PATH: Optional[str] = None
LOCAL_RANK: int = 0
CUDA_VISIBLE_DEVICES: str | None = None
CUDA_VISIBLE_DEVICES: Optional[str] = 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: str | None = None
FASTVIDEO_LOGGING_CONFIG_PATH: Optional[str] = None
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_ATTENTION_CONFIG: str | None = None
FASTVIDEO_ATTENTION_BACKEND: Optional[str] = None
FASTVIDEO_ATTENTION_CONFIG: Optional[str] = None
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: str | None = None
NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: str | None = None
MAX_JOBS: Optional[str] = None
NVCC_THREADS: Optional[str] = None
CMAKE_BUILD_TYPE: Optional[str] = None
VERBOSE: bool = False
FASTVIDEO_SERVER_DEV_MODE: bool = False
@@ -43,7 +42,7 @@ def get_default_config_root() -> str:
)
def maybe_convert_int(value: str | None) -> int | None:
def maybe_convert_int(value: Optional[str]) -> Optional[int]:
if value is None:
return None
return int(value)
@@ -54,7 +53,7 @@ def maybe_convert_int(value: str | None) -> int | None:
# begin-env-vars-definition
environment_variables: dict[str, Callable[[], Any]] = {
environment_variables: Dict[str, Callable[[], Any]] = {
# ================== Installation Time Env Vars ==================
@@ -0,0 +1,30 @@
from fastvideo import VideoGenerator
# from fastvideo.v1.configs.sample import SamplingParam
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"FastVideo/FastHunyuan-diffusers",
# if num_gpus > 1, FastVideo will automatically handle distributed setup
num_gpus=4,
)
# sampling_param = SamplingParam.from_pretrained("/workspace/data/Wan-AI/Wan2.1-I2V-14B-480P-Diffusers")
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = "A beautiful woman in a red dress walking down a street"
video = generator.generate_video(prompt)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = "A beautiful woman in a blue dress walking down a street"
video2 = generator.generate_video(prompt2)
if __name__ == "__main__":
main()
+16 -17
View File
@@ -4,10 +4,9 @@
import argparse
import dataclasses
from collections.abc import Callable
from contextlib import contextmanager
from dataclasses import field
from typing import Any
from typing import Any, Callable, List, Optional, Tuple
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.logger import init_logger
@@ -39,17 +38,17 @@ class FastVideoArgs:
# HuggingFace specific parameters
trust_remote_code: bool = False
revision: str | None = None
revision: Optional[str] = None
# Parallelism
num_gpus: int = 1
tp_size: int | None = None
sp_size: int | None = None
dist_timeout: int | None = None # timeout for torch.distributed
tp_size: Optional[int] = None
sp_size: Optional[int] = None
dist_timeout: Optional[int] = None # timeout for torch.distributed
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: float | None = None
flow_shift: Optional[float] = None
output_type: str = "pil"
@@ -73,32 +72,32 @@ class FastVideoArgs:
"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: str | None = None
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
use_cpu_offload: bool = False
disable_autocast: bool = False
# StepVideo specific parameters
pos_magic: str | None = None
neg_magic: str | None = None
timesteps_scale: bool | None = None
pos_magic: Optional[str] = None
neg_magic: Optional[str] = None
timesteps_scale: Optional[bool] = None
# Logging
log_level: str = "info"
# Inference parameters
device_str: str | None = None
device_str: Optional[str] = None
device = None
def __post_init__(self):
@@ -377,7 +376,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.
+5 -5
View File
@@ -5,7 +5,7 @@ import time
from collections import defaultdict
from contextlib import contextmanager
from dataclasses import dataclass
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Optional
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: ForwardBatch | None = None
forward_batch: Optional[ForwardBatch] = None
_forward_context: ForwardContext | None = None
_forward_context: Optional[ForwardContext] = 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: ForwardBatch | None = None,
fastvideo_args: FastVideoArgs | None = None):
forward_batch: Optional[ForwardBatch] = None,
fastvideo_args: Optional[FastVideoArgs] = None):
"""A context manager that stores the current forward context,
can be attention metadata, etc.
Here we can inject common logic for every model forward pass.
+3 -3
View File
@@ -8,7 +8,7 @@ This module provides classes and functions for running inference with diffusion
"""
import time
from typing import Any
from typing import Any, Dict
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
+2 -3
View File
@@ -1,8 +1,7 @@
# 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 collections.abc import Callable
from typing import Any
from typing import Any, Callable, Dict, Type
import torch.nn as nn
@@ -82,7 +81,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
+5 -4
View File
@@ -1,6 +1,7 @@
# 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
@@ -21,7 +22,7 @@ class RMSNorm(CustomOp):
hidden_size: int,
eps: float = 1e-6,
dtype: torch.dtype = torch.float32,
var_hidden_size: int | None = None,
var_hidden_size: Optional[int] = None,
has_weight: bool = True,
) -> None:
super().__init__()
@@ -39,8 +40,8 @@ class RMSNorm(CustomOp):
def forward_native(
self,
x: torch.Tensor,
residual: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""PyTorch-native implementation equivalent to forward()."""
orig_dtype = x.dtype
x = x.to(torch.float32)
@@ -129,7 +130,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.
+42 -34
View File
@@ -2,6 +2,7 @@
# 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
@@ -39,7 +40,7 @@ WEIGHT_LOADER_V2_SUPPORTED = [
def adjust_scalar_to_fused_array(
param: torch.Tensor, loaded_weight: torch.Tensor,
shard_id: str | int) -> tuple[torch.Tensor, torch.Tensor]:
shard_id: Union[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
@@ -90,7 +91,7 @@ class LinearMethodBase(QuantizeMethodBase):
def apply(self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None) -> torch.Tensor:
bias: Optional[torch.Tensor] = 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
@@ -115,7 +116,7 @@ class UnquantizedLinearMethod(LinearMethodBase):
def apply(self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None) -> torch.Tensor:
bias: Optional[torch.Tensor] = None) -> torch.Tensor:
return F.linear(x, layer.weight, bias)
@@ -137,8 +138,8 @@ class LinearBase(torch.nn.Module):
input_size: int,
output_size: int,
skip_bias_add: bool = False,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
@@ -151,13 +152,14 @@ class LinearBase(torch.nn.Module):
params_dtype = torch.get_default_dtype()
self.params_dtype = params_dtype
if quant_config is None:
self.quant_method: QuantizeMethodBase | None = UnquantizedLinearMethod(
)
self.quant_method: Optional[
QuantizeMethodBase] = UnquantizedLinearMethod()
else:
self.quant_method = quant_config.get_quant_method(self,
prefix=prefix)
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
def forward(self,
x: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
raise NotImplementedError
@@ -180,8 +182,8 @@ class ReplicatedLinear(LinearBase):
output_size: int,
bias: bool = True,
skip_bias_add: bool = False,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
super().__init__(input_size,
output_size,
@@ -221,7 +223,8 @@ 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, Parameter | None]:
def forward(self,
x: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
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)
@@ -265,9 +268,9 @@ class ColumnParallelLinear(LinearBase):
bias: bool = True,
gather_output: bool = False,
skip_bias_add: bool = False,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
output_sizes: list[int] | None = None,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
output_sizes: Optional[list[int]] = None,
prefix: str = ""):
# Divide the weight matrix along the last dimension.
self.tp_size = get_tensor_model_parallel_world_size()
@@ -342,8 +345,9 @@ 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, Parameter | None]:
def forward(
self,
input_: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
bias = self.bias if not self.skip_bias_add else None
# Matrix multiply.
@@ -395,8 +399,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
bias: bool = True,
gather_output: bool = False,
skip_bias_add: bool = False,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
self.output_sizes = output_sizes
tp_size = get_tensor_model_parallel_world_size()
@@ -413,7 +417,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
def weight_loader(self,
param: Parameter,
loaded_weight: torch.Tensor,
loaded_shard_id: int | None = None) -> None:
loaded_shard_id: Optional[int] = None) -> None:
param_data = param.data
output_dim = getattr(param, "output_dim", None)
@@ -506,8 +510,10 @@ 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)
@@ -519,7 +525,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
def weight_loader_v2(self,
param: BasevLLMParameter,
loaded_weight: torch.Tensor,
loaded_shard_id: int | None = None) -> None:
loaded_shard_id: Optional[int] = None) -> None:
if loaded_shard_id is None:
if isinstance(param, PerTensorScaleParameter):
param.load_merged_column_weight(loaded_weight=loaded_weight,
@@ -592,11 +598,11 @@ class QKVParallelLinear(ColumnParallelLinear):
hidden_size: int,
head_size: int,
total_num_heads: int,
total_num_kv_heads: int | None = None,
total_num_kv_heads: Optional[int] = None,
bias: bool = True,
skip_bias_add: bool = False,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
self.hidden_size = hidden_size
self.head_size = head_size
@@ -631,7 +637,7 @@ class QKVParallelLinear(ColumnParallelLinear):
quant_config=quant_config,
prefix=prefix)
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> int | None:
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> Optional[int]:
shard_offset_mapping = {
"q": 0,
"k": self.num_heads * self.head_size,
@@ -640,7 +646,7 @@ class QKVParallelLinear(ColumnParallelLinear):
}
return shard_offset_mapping.get(loaded_shard_id)
def _get_shard_size_mapping(self, loaded_shard_id: str) -> int | None:
def _get_shard_size_mapping(self, loaded_shard_id: str) -> Optional[int]:
shard_size_mapping = {
"q": self.num_heads * self.head_size,
"k": self.num_kv_heads * self.head_size,
@@ -673,8 +679,10 @@ 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)
@@ -686,7 +694,7 @@ class QKVParallelLinear(ColumnParallelLinear):
def weight_loader_v2(self,
param: BasevLLMParameter,
loaded_weight: torch.Tensor,
loaded_shard_id: str | None = None):
loaded_shard_id: Optional[str] = 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)
@@ -712,7 +720,7 @@ class QKVParallelLinear(ColumnParallelLinear):
def weight_loader(self,
param: Parameter,
loaded_weight: torch.Tensor,
loaded_shard_id: str | None = None):
loaded_shard_id: Optional[str] = None):
param_data = param.data
output_dim = getattr(param, "output_dim", None)
@@ -837,9 +845,9 @@ class RowParallelLinear(LinearBase):
bias: bool = True,
input_is_parallel: bool = True,
skip_bias_add: bool = False,
params_dtype: torch.dtype | None = None,
params_dtype: Optional[torch.dtype] = None,
reduce_results: bool = True,
quant_config: QuantizationConfig | None = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
# Divide the weight matrix along the first dimension.
self.tp_rank = get_tensor_model_parallel_rank()
@@ -913,7 +921,7 @@ class RowParallelLinear(LinearBase):
param.load_row_parallel_weight(loaded_weight=loaded_weight)
def forward(self, input_) -> tuple[torch.Tensor, Parameter | None]:
def forward(self, input_) -> tuple[torch.Tensor, Optional[Parameter]]:
if self.input_is_parallel:
input_parallel = input_
else:
+4 -2
View File
@@ -1,5 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional
import torch
import torch.nn as nn
@@ -16,10 +18,10 @@ class MLP(nn.Module):
self,
input_dim: int,
mlp_hidden_dim: int,
output_dim: int | None = None,
output_dim: Optional[int] = None,
bias: bool = True,
act_type: str = "gelu_pytorch_tanh",
dtype: torch.dtype | None = None,
dtype: Optional[torch.dtype] = None,
prefix: str = "",
):
super().__init__()
@@ -3,7 +3,7 @@
import inspect
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, Optional
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) -> QuantizationMethods | None:
def override_quantization_method(
cls, hf_quant_cfg, user_quant) -> Optional[QuantizationMethods]:
"""
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) -> QuantizeMethodBase | None:
prefix: str) -> Optional[QuantizeMethodBase]:
"""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) -> str | None:
return None
def get_cache_scale(self, name: str) -> Optional[str]:
return None
+20 -20
View File
@@ -23,7 +23,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
"""Rotary Positional Embeddings."""
from typing import Any
from typing import Any, Dict, List, Optional, Tuple, Union
import torch
@@ -84,7 +84,7 @@ class RotaryEmbedding(CustomOp):
head_size: int,
rotary_dim: int,
max_position_embeddings: int,
base: int | float,
base: Union[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: int | float) -> torch.Tensor:
def _compute_inv_freq(self, base: Union[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: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
offsets: Optional[torch.Tensor] = 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: int | tuple[int, ...], dim: int = 2) -> tuple[int, ...]:
def _to_tuple(x: Union[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: int | tuple[int, ...], dim: int = 2) -> tuple[int, ...]:
raise ValueError(f"Expected length {dim} or int, but got {x}")
def get_meshgrid_nd(start: int | tuple[int, ...],
*args: int | tuple[int, ...],
def get_meshgrid_nd(start: Union[int, Tuple[int, ...]],
*args: Union[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: int | tuple[int, ...],
def get_1d_rotary_pos_embed(
dim: int,
pos: torch.FloatTensor | int,
pos: Union[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: float | list[float] = 1.0,
interpolation_factor: float | list[float] = 1.0,
theta_rescale_factor: Union[float, List[float]] = 1.0,
interpolation_factor: Union[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: int | float,
base: Union[int, float],
is_neox_style: bool = True,
rope_scaling: dict[str, Any] | None = None,
dtype: torch.dtype | None = None,
rope_scaling: Optional[Dict[str, Any]] = None,
dtype: Optional[torch.dtype] = None,
partial_rotary_factor: float = 1.0,
) -> RotaryEmbedding:
if dtype is None:
+2 -1
View File
@@ -1,6 +1,7 @@
# 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
@@ -9,7 +10,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),
+3 -2
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Optional
import torch
import torch.nn as nn
@@ -35,7 +36,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:
@@ -132,7 +133,7 @@ class ModulateProjection(nn.Module):
hidden_size: int,
factor: int = 2,
act_layer: str = "silu",
dtype: torch.dtype | None = None,
dtype: Optional[torch.dtype] = None,
prefix: str = "",
):
super().__init__()
+11 -11
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Sequence
from dataclasses import dataclass
from typing import List, Optional, Sequence, Tuple
import torch
import torch.nn.functional as F
@@ -24,7 +24,7 @@ class UnquantizedEmbeddingMethod(QuantizeMethodBase):
def create_weights(self, layer: torch.nn.Module,
input_size_per_partition: int,
output_partition_sizes: list[int], input_size: int,
output_partition_sizes: List[int], input_size: int,
output_size: int, params_dtype: torch.dtype,
**extra_weight_attrs):
"""Create weights for embedding layer."""
@@ -39,7 +39,7 @@ class UnquantizedEmbeddingMethod(QuantizeMethodBase):
def apply(self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None) -> torch.Tensor:
bias: Optional[torch.Tensor] = 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: torch.dtype | None = None,
org_num_embeddings: int | None = None,
params_dtype: Optional[torch.dtype] = None,
org_num_embeddings: Optional[int] = None,
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
quant_config: QuantizationConfig | None = None,
quant_config: Optional[QuantizationConfig] = 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) -> list[int] | None:
def get_sharded_to_full_mapping(self) -> Optional[List[int]]:
"""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,
+3 -2
View File
@@ -11,7 +11,7 @@ from logging import Logger
from logging.config import dictConfig
from os import path
from types import MethodType
from typing import Any, cast
from typing import Any, Optional, cast
import fastvideo.v1.envs as envs
@@ -278,7 +278,8 @@ 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: str | None = None):
def enable_trace_function_call(log_file_path: str,
root_dir: Optional[str] = None):
"""
Enable tracing of every function call in code under `root_dir`.
This is useful for debugging hangs or crashes.
+7 -7
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import Any
from typing import Any, List, Optional, Tuple, Union
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:
@@ -44,10 +44,10 @@ class BaseDiT(nn.Module, ABC):
@abstractmethod
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs) -> torch.Tensor:
pass
@@ -63,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
@@ -81,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:
+560
View File
@@ -0,0 +1,560 @@
from typing import List, Optional, Tuple, Union
import torch
import torch.nn as nn
from diffusers.models.normalization import RMSNorm
from fastvideo.v1.attention import DistributedAttention
from fastvideo.v1.configs.models.dits import FluxImageConfig
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.v1.layers.linear import ReplicatedLinear
from fastvideo.v1.layers.mlp import MLP
from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
TimestepEmbedder)
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.platforms import _Backend
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal DiT block with separate modulation for text and image/video,
using distributed attention and linear layers.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
prefix: str = "",
):
super().__init__()
self.deterministic = False
self.num_attention_heads = num_attention_heads
head_dim = hidden_size // num_attention_heads
mlp_hidden_dim = int(hidden_size * mlp_ratio)
# Image modulation components
self.img_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.img_mod",
)
# Fused operations for image stream
self.img_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_mlp_residual = ScaleResidual()
# Image attention components
self.img_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_qkv")
self.img_attn_q_norm = RMSNorm(head_dim, eps=1e-6)
self.img_attn_k_norm = RMSNorm(head_dim, eps=1e-6)
self.img_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_proj")
self.img_mlp = MLP(hidden_size,
mlp_hidden_dim,
bias=True,
dtype=dtype,
prefix=f"{prefix}.img_mlp")
# Text modulation components
self.txt_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.txt_mod",
)
# Fused operations for text stream
self.txt_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_mlp_residual = ScaleResidual()
# Text attention components
self.txt_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype)
# QK norm layers for text
self.txt_attn_q_norm = RMSNorm(head_dim, eps=1e-6)
self.txt_attn_k_norm = RMSNorm(head_dim, eps=1e-6)
self.txt_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype)
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
# Distributed attention
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn")
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
vec: torch.Tensor,
freqs_cis_img: Tuple[torch.Tensor, torch.Tensor],
) -> Tuple[torch.Tensor, torch.Tensor]:
# Process modulation vectors
img_mod_outputs = self.img_mod(vec)
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=-1)
txt_mod_outputs = self.txt_mod(vec)
(
txt_attn_shift,
txt_attn_scale,
txt_attn_gate,
txt_mlp_shift,
txt_mlp_scale,
txt_mlp_gate,
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
# Prepare image for attention using fused operation
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale)
# Get QKV for image
img_qkv, _ = self.img_attn_qkv(img_attn_input)
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
# Split QKV
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
self.num_attention_heads, -1)
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
2]
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply rotary embeddings for image
cos, sin = freqs_cis_img
img_q, img_k = _apply_rotary_emb(
img_q, cos, sin,
is_neox_style=False), _apply_rotary_emb(img_k,
cos,
sin,
is_neox_style=False)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
# Get QKV for text
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
# Split QKV
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
self.num_attention_heads, -1)
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
2]
# Apply QK-Norm if needed
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
# Run distributed attention
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v)
img_attn_out, _ = self.img_attn_proj(
img_attn.view(batch_size, image_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale)
# Process image MLP
img_mlp_out = self.img_mlp(img_mlp_input)
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
# Process text attention output
txt_attn_out, _ = self.txt_attn_proj(
txt_attn.reshape(batch_size, text_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale)
# Process text MLP
txt_mlp_out = self.txt_mlp(txt_mlp_input)
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
return img, txt
class MMSingleStreamBlock(nn.Module):
"""
A DiT block with parallel linear layers using distributed attention
and tensor parallelism.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
prefix: str = "",
):
super().__init__()
self.deterministic = False
self.hidden_size = hidden_size
self.num_attention_heads = num_attention_heads
head_dim = hidden_size // num_attention_heads
mlp_hidden_dim = int(hidden_size * mlp_ratio)
self.mlp_hidden_dim = mlp_hidden_dim
# Combined QKV and MLP input projection
self.linear1 = ReplicatedLinear(hidden_size,
hidden_size * 3 + mlp_hidden_dim,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear1")
# Combined projection and MLP output
self.linear2 = ReplicatedLinear(hidden_size + mlp_hidden_dim,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear2")
# QK norm layers
self.q_norm = RMSNorm(head_dim, eps=1e-6)
self.k_norm = RMSNorm(head_dim, eps=1e-6)
# Fused operations with better naming
self.input_norm_scale_shift = LayerNormScaleShift(
hidden_size,
norm_type="layer",
eps=1e-6,
elementwise_affine=False,
dtype=dtype)
self.output_residual = ScaleResidual()
# Activation function
self.mlp_act = nn.GELU(approximate="tanh")
# Modulation
self.modulation = ModulateProjection(hidden_size,
factor=3,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.modulation")
# Distributed attention
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn")
def forward(
self,
x: torch.Tensor,
vec: torch.Tensor,
txt_len: int,
freqs_cis_img: Tuple[torch.Tensor, torch.Tensor],
) -> torch.Tensor:
# Process modulation
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
# Apply pre-norm and modulation using fused operation
x_mod = self.input_norm_scale_shift(x, mod_shift, mod_scale)
# Get combined projections
linear1_out, _ = self.linear1(x_mod)
# Split into QKV and MLP parts
qkv, mlp = torch.split(linear1_out,
[3 * self.hidden_size, self.mlp_hidden_dim],
dim=-1)
# Process QKV
batch_size, seq_len = qkv.shape[0], qkv.shape[1]
qkv = qkv.view(batch_size, seq_len, 3, self.num_attention_heads, -1)
q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2]
# Apply QK-Norm
q = self.q_norm(q).to(v.dtype)
k = self.k_norm(k).to(v.dtype)
# Split into image and text parts
img_q, txt_q = q[:, :-txt_len], q[:, -txt_len:]
img_k, txt_k = k[:, :-txt_len], k[:, -txt_len:]
img_v, txt_v = v[:, :-txt_len], v[:, -txt_len:]
# Apply rotary embeddings to image parts
cos, sin = freqs_cis_img
img_q, img_k = _apply_rotary_emb(
img_q, cos, sin,
is_neox_style=False), _apply_rotary_emb(img_k,
cos,
sin,
is_neox_style=False)
# Run distributed attention
img_attn_output, txt_attn_output = self.attn(img_q, img_k, img_v, txt_q,
txt_k, txt_v)
attn_output = torch.cat((img_attn_output, txt_attn_output),
dim=1).view(batch_size, seq_len, -1)
# Process MLP activation
mlp_output = self.mlp_act(mlp)
# Combine attention and MLP outputs
combined = torch.cat((attn_output, mlp_output), dim=-1)
# Final projection
output, _ = self.linear2(combined)
# Apply residual connection with gating using fused operation
return self.output_residual(x, output, mod_gate)
class FinalLayer(nn.Module):
"""
The final layer of DiT that projects features to pixel space.
"""
def __init__(self,
hidden_size,
patch_size,
out_channels,
dtype=None,
prefix: str = "") -> None:
super().__init__()
# Normalization
self.norm_final = nn.LayerNorm(hidden_size,
eps=1e-6,
elementwise_affine=False,
dtype=dtype)
output_dim = patch_size**3 * out_channels
self.linear = ReplicatedLinear(hidden_size,
output_dim,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear")
# Modulation
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
def forward(self, img, vec):
scale, shift = self.adaLN_modulation(vec).chunk(2, dim=-1)
img = self.norm_final(img) * (1.0 +
scale.unsqueeze(1)) + shift.unsqueeze(1)
img, _ = self.linear(img)
return img
class FluxTransformer2DModel(BaseDiT):
_fsdp_shard_conditions = FluxImageConfig()._fsdp_shard_conditions
_supported_attention_backends = FluxImageConfig(
)._supported_attention_backends
_param_names_mapping = FluxImageConfig()._param_names_mapping
def __init__(self, config: FluxImageConfig) -> None:
super().__init__(config=config)
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.num_channels_latents = config.num_channels_latents
self.text_states_dim = config.joint_attention_dim
self.text_states_dim_2 = config.pooled_projection_dim
self.rope_dim_list = list(config.axes_dims_rope)
self.rope_theta = config.rope_theta
self.out_channels = config.out_channels
self.patch_size = config.patch_size
self.img_in = ReplicatedLinear(config.in_channels,
self.hidden_size,
params_dtype=config.dtype,
prefix=f"{config.prefix}.img_in")
self.txt_in = ReplicatedLinear(self.text_states_dim,
self.hidden_size,
params_dtype=config.dtype,
prefix=f"{config.prefix}.txt_in")
self.time_in = TimestepEmbedder(self.hidden_size,
act_layer="silu",
dtype=config.dtype,
prefix=f"{config.prefix}.time_in")
self.txt2_in = MLP(self.text_states_dim_2,
self.hidden_size,
self.hidden_size,
act_type="silu",
dtype=config.dtype,
prefix=f"{config.prefix}.txt2_in")
self.guidance_in = (TimestepEmbedder(
self.hidden_size,
act_layer="silu",
dtype=config.dtype,
prefix=f"{config.prefix}.guidance_in")
if config.guidance_embeds else None)
# Double blocks
self.double_blocks = nn.ModuleList([
MMDoubleStreamBlock(
self.hidden_size,
self.num_attention_heads,
dtype=config.dtype,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.double_blocks.{i}")
for i in range(config.num_layers)
])
# Single blocks
self.single_blocks = nn.ModuleList([
MMSingleStreamBlock(
self.hidden_size,
self.num_attention_heads,
dtype=config.dtype,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.single_blocks.{i+config.num_layers}")
for i in range(config.num_single_layers)
])
self.final_layer = FinalLayer(config.hidden_size,
self.patch_size,
self.out_channels,
dtype=config.dtype,
prefix=f"{config.prefix}.final_layer")
self.__post_init__()
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs):
"""
Forward pass of the FluxTransformer2DModel.
Args:
hidden_states: Input image latents [B, N, C]
encoder_hidden_states: Text embeddings [B, L, D]
timestep: Diffusion timestep
guidance: Guidance scale for CFG
Returns:
Tuple of (output)
"""
h = kwargs.pop("height_latents") or None
w = kwargs.pop("width_latents") or None
assert h is not None and w is not None
img = x = hidden_states
# Match diffusers implementation by multiplying timestep by 1000
t = timestep.to(img.dtype)
# Split text embeddings - first token is global, rest are per-token
txt = encoder_hidden_states[1]
text_states_2 = encoder_hidden_states[0]
# Get spatial dimensions
# _, _, oh, ow = img.shape
th, tw = (h // self.patch_size // 2, w // self.patch_size // 2)
# Get rotary embeddings
freqs_cos_img, freqs_sin_img = get_rotary_pos_embed(
(1, th, tw),
self.hidden_size,
self.num_attention_heads,
self.rope_dim_list,
self.rope_theta,
shard_dim=1,
)
freqs_cos_img = freqs_cos_img.to(img.device)
freqs_sin_img = freqs_sin_img.to(img.device)
freqs_cis_img = (freqs_cos_img, freqs_sin_img)
# Prepare modulation vectors
vec = self.time_in(t)
# Add text modulation
vec = vec + self.txt2_in(text_states_2)
# Add guidance modulation
if self.guidance_in is not None and guidance is not None:
vec = vec + self.guidance_in(guidance)
# embed text and image
img, _ = self.img_in(img)
txt, _ = self.txt_in(txt)
img_seq_len = img.shape[1]
txt_seq_len = txt.shape[1]
# Process through double stream blocks
for index, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis_img]
img, txt = block(*double_block_args)
# Merge txt and img to pass through single stream blocks
x = torch.cat((img, txt), 1)
# Process through single stream blocks
for index, block in enumerate(self.single_blocks):
single_block_args = [x, vec, txt_seq_len, freqs_cis_img]
x = block(*single_block_args)
# Extract image features
img = x[:, :img_seq_len, ...]
# Final layer
img = self.final_layer(img, vec)
return img
+11 -9
View File
@@ -1,5 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from typing import List, Optional, Tuple, Union
import numpy as np
import torch
import torch.nn as nn
@@ -94,8 +96,8 @@ class MMDoubleStreamBlock(nn.Module):
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[_Backend, ...] | None = None,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
prefix: str = "",
):
super().__init__()
@@ -200,7 +202,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)
(
@@ -301,8 +303,8 @@ class MMSingleStreamBlock(nn.Module):
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[_Backend, ...] | None = None,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
prefix: str = "",
):
super().__init__()
@@ -364,7 +366,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)
@@ -540,10 +542,10 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
# TODO: change output to a dict
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs):
"""
+16 -17
View File
@@ -10,6 +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 Dict, Optional, Tuple
import torch
from einops import rearrange, repeat
@@ -54,7 +55,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:
@@ -142,7 +143,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,
@@ -189,10 +190,8 @@ class SelfAttention(nn.Module):
outs = []
idx = 0
for (chunk_size, cos_i, sin_i) in zip(self.rope_split,
cos_splits,
sin_splits,
strict=False):
for (chunk_size, cos_i, sin_i) in zip(self.rope_split, cos_splits,
sin_splits):
# slice the corresponding channels
x_chunk = x[..., idx:idx + chunk_size] # [B,S,H,chunk_size]
idx += chunk_size
@@ -332,8 +331,8 @@ class AdaLayerNormSingle(nn.Module):
def forward(
self,
timestep: torch.Tensor,
added_cond_kwargs: dict[str, torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
embedded_timestep = self.emb(timestep * self.time_step_rescale)
out, _ = self.linear(self.silu(embedded_timestep))
@@ -378,7 +377,7 @@ class StepVideoTransformerBlock(nn.Module):
dim: int,
attention_head_dim: int,
norm_eps: float = 1e-5,
ff_inner_dim: int | None = None,
ff_inner_dim: Optional[int] = None,
ff_bias: bool = False,
attention_type: str = 'torch'):
super().__init__()
@@ -418,7 +417,7 @@ class StepVideoTransformerBlock(nn.Module):
kv: torch.Tensor,
t_expand: torch.LongTensor,
attn_mask=None,
rope_positions: list | None = None,
rope_positions: Optional[list] = None,
cos_sin=None,
mask_strategy=None) -> torch.Tensor:
@@ -539,7 +538,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 +593,12 @@ class StepVideoModel(BaseDiT):
def forward(
self,
hidden_states: torch.Tensor,
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,
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,
return_dict: bool = True,
mask_strategy=None,
guidance=None,
+11 -10
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import List, Optional, Tuple, Union
import numpy as np
import torch
@@ -52,7 +53,7 @@ class WanTimeTextImageEmbedding(nn.Module):
dim: int,
time_freq_dim: int,
text_embed_dim: int,
image_embed_dim: int | None = None,
image_embed_dim: Optional[int] = None,
):
super().__init__()
@@ -75,7 +76,7 @@ class WanTimeTextImageEmbedding(nn.Module):
self,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: torch.Tensor | None = None,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
):
temb = self.time_embedder(timestep)
timestep_proj = self.time_modulation(temb)
@@ -172,7 +173,7 @@ class WanI2VCrossAttention(WanSelfAttention):
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
supported_attention_backends: tuple[_Backend, ...] | None = None
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps,
supported_attention_backends)
@@ -221,9 +222,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: int | None = None,
supported_attention_backends: tuple[_Backend, ...]
| None = None,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = ""):
super().__init__()
@@ -291,7 +292,7 @@ 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)
@@ -416,10 +417,10 @@ class WanTransformer3DModel(CachableDiT):
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs) -> torch.Tensor:
forward_batch = get_forward_context().forward_batch
+10 -9
View File
@@ -1,4 +1,5 @@
from abc import ABC, abstractmethod
from typing import Optional, Tuple
import torch
from torch import nn
@@ -10,7 +11,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:
@@ -23,21 +24,21 @@ class TextEncoder(nn.Module, ABC):
@abstractmethod
def forward(self,
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,
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,
**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:
@@ -54,5 +55,5 @@ class ImageEncoder(nn.Module, ABC):
pass
@property
def supported_attention_backends(self) -> tuple[_Backend, ...]:
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
return self._supported_attention_backends
+38 -38
View File
@@ -3,7 +3,7 @@
# Adapted from transformers: https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py
"""Minimal implementation of CLIPVisionModel intended to be only used
within a vision language model."""
from collections.abc import Iterable
from typing import Iterable, Optional, Set, Tuple, Union
import torch
import torch.nn as nn
@@ -91,9 +91,9 @@ class CLIPTextEmbeddings(nn.Module):
def forward(
self,
input_ids: torch.LongTensor | None = None,
position_ids: torch.LongTensor | None = None,
inputs_embeds: torch.FloatTensor | None = None,
input_ids: Optional[torch.LongTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = 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: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
@@ -209,8 +209,8 @@ class CLIPMLP(nn.Module):
def __init__(
self,
config: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -239,8 +239,8 @@ class CLIPEncoderLayer(nn.Module):
def __init__(
self,
config: CLIPTextConfig | CLIPVisionConfig,
quant_config: QuantizationConfig | None = None,
config: Union[CLIPTextConfig, CLIPVisionConfig],
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -284,9 +284,9 @@ class CLIPEncoder(nn.Module):
def __init__(
self,
config: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = None,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = 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
) -> torch.Tensor | list[torch.Tensor]:
self, inputs_embeds: torch.Tensor, return_all_hidden_states: bool
) -> Union[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: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = None,
quant_config: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = None,
prefix: str = ""):
super().__init__()
self.config = config
@@ -348,11 +348,11 @@ class CLIPTextTransformer(nn.Module):
def forward(
self,
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,
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,
) -> BaseEncoderOutput:
r"""
Returns:
@@ -440,11 +440,11 @@ class CLIPTextModel(TextEncoder):
def forward(
self,
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,
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,
**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: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = None,
require_post_norm: bool | None = None,
quant_config: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = None,
require_post_norm: Optional[bool] = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -540,7 +540,7 @@ class CLIPVisionTransformer(nn.Module):
def forward(
self,
pixel_values: torch.Tensor,
feature_sample_layers: list[int] | None = None,
feature_sample_layers: Optional[list[int]] = 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: list[int] | None = None,
feature_sample_layers: Optional[list[int]] = 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:
+17 -18
View File
@@ -23,8 +23,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
"""Inference-only LLaMA model compatible with HuggingFace weights."""
from collections.abc import Iterable
from typing import Any
from typing import Any, Dict, Iterable, Optional, Set, Tuple
import torch
from torch import nn
@@ -53,7 +52,7 @@ class LlamaMLP(nn.Module):
hidden_size: int,
intermediate_size: int,
hidden_act: str,
quant_config: QuantizationConfig | None = None,
quant_config: Optional[QuantizationConfig] = None,
bias: bool = False,
prefix: str = "",
) -> None:
@@ -93,9 +92,9 @@ class LlamaAttention(nn.Module):
num_heads: int,
num_kv_heads: int,
rope_theta: float = 10000,
rope_scaling: dict[str, Any] | None = None,
rope_scaling: Optional[Dict[str, Any]] = None,
max_position_embeddings: int = 8192,
quant_config: QuantizationConfig | None = None,
quant_config: Optional[QuantizationConfig] = None,
bias: bool = False,
bias_o_proj: bool = False,
prefix: str = "") -> None:
@@ -202,7 +201,7 @@ class LlamaDecoderLayer(nn.Module):
def __init__(
self,
config: LlamaConfig,
quant_config: QuantizationConfig | None = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -255,8 +254,8 @@ class LlamaDecoderLayer(nn.Module):
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
residual: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
residual: Optional[torch.Tensor],
) -> Tuple[torch.Tensor, torch.Tensor]:
# Self Attention
if residual is None:
residual = hidden_states
@@ -319,11 +318,11 @@ class LlamaModel(TextEncoder):
def forward(
self,
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,
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,
**kwargs,
) -> BaseEncoderOutput:
output_hidden_states = (output_hidden_states
@@ -340,7 +339,7 @@ class LlamaModel(TextEncoder):
0, hidden_states.shape[1],
device=hidden_states.device).unsqueeze(0)
all_hidden_states: tuple[Any, ...] | None = (
all_hidden_states: Optional[Tuple[Any, ...]] = (
) if output_hidden_states else None
for layer in self.layers:
if all_hidden_states is not None:
@@ -368,8 +367,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"),
@@ -379,7 +378,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
@@ -401,7 +400,7 @@ class LlamaModel(TextEncoder):
# continue
if "scale" in name:
# Remapping the name of FP8 kv-scale.
kv_scale_name: str | None = maybe_remap_kv_scale_name(
kv_scale_name: Optional[str] = maybe_remap_kv_scale_name(
name, params_dict)
if kv_scale_name is None:
continue
+9 -8
View File
@@ -13,6 +13,7 @@
# ==============================================================================
import os
from functools import wraps
from typing import List, Optional
import torch
import torch.nn as nn
@@ -178,10 +179,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)
@@ -346,9 +347,9 @@ class MultiQueryAttention(nn.Module):
def forward(
self,
x: torch.Tensor,
mask: torch.Tensor | None,
cu_seqlens: torch.Tensor | None,
max_seq_len: torch.Tensor | None,
mask: Optional[torch.Tensor],
cu_seqlens: Optional[torch.Tensor],
max_seq_len: Optional[torch.Tensor],
):
seqlen, bsz, dim = x.shape
xqkv = self.wqkv(x)
@@ -470,9 +471,9 @@ class TransformerBlock(nn.Module):
def forward(
self,
x: torch.Tensor,
mask: torch.Tensor | None,
cu_seqlens: torch.Tensor | None,
max_seq_len: torch.Tensor | None,
mask: Optional[torch.Tensor],
cu_seqlens: Optional[torch.Tensor],
max_seq_len: Optional[torch.Tensor],
):
residual = self.attention.forward(self.attention_norm(x), mask,
cu_seqlens, max_seq_len)
+29 -29
View File
@@ -20,8 +20,8 @@
"""PyTorch T5 & UMT5 model."""
import math
from collections.abc import Iterable
from dataclasses import dataclass
from typing import Iterable, Optional, Set, Tuple
import torch
import torch.nn.functional as F
@@ -64,7 +64,7 @@ class T5DenseActDense(nn.Module):
def __init__(self,
config: T5Config,
quant_config: QuantizationConfig | None = None):
quant_config: Optional[QuantizationConfig] = 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: QuantizationConfig | None = None):
quant_config: Optional[QuantizationConfig] = 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: QuantizationConfig | None = None):
quant_config: Optional[QuantizationConfig] = 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: QuantizationConfig | None = None,
quant_config: Optional[QuantizationConfig] = 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: AttentionMetadata | None = None,
attn_metadata: Optional[AttentionMetadata] = 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: QuantizationConfig | None = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
@@ -361,7 +361,7 @@ class T5LayerSelfAttention(nn.Module):
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
attn_metadata: AttentionMetadata | None = None,
attn_metadata: Optional[AttentionMetadata] = 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: QuantizationConfig | None = None,
quant_config: Optional[QuantizationConfig] = 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: AttentionMetadata | None = None,
attn_metadata: Optional[AttentionMetadata] = 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: QuantizationConfig | None = None,
quant_config: Optional[QuantizationConfig] = 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: AttentionMetadata | None = None,
attn_metadata: Optional[AttentionMetadata] = 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: QuantizationConfig | None = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
is_umt5: bool = False):
super().__init__()
@@ -524,11 +524,11 @@ class T5EncoderModel(TextEncoder):
def forward(
self,
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,
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,
**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: 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,
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,
**kwargs,
) -> BaseEncoderOutput:
attn_metadata = AttentionMetadata(None)
@@ -630,8 +630,8 @@ class UMT5EncoderModel(TextEncoder):
attention_mask=attention_mask,
)
def load_weights(self, weights: Iterable[tuple[str,
torch.Tensor]]) -> set[str]:
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
(".qkv_proj", ".q", "q"),
@@ -639,7 +639,7 @@ class UMT5EncoderModel(TextEncoder):
(".qkv_proj", ".v", "v"),
]
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
loaded_params: Set[str] = set()
for name, loaded_weight in weights:
loaded = False
if "decoder" in name or "lm_head" in name:
+4 -4
View File
@@ -2,7 +2,7 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/models/vision.py
from abc import ABC, abstractmethod
from typing import Generic, TypeVar
from typing import Generic, Optional, TypeVar, Union
import torch
from transformers import PretrainedConfig
@@ -48,9 +48,9 @@ class VisionEncoderInfo(ABC, Generic[_C]):
def resolve_visual_encoder_outputs(
encoder_outputs: torch.Tensor | list[torch.Tensor],
feature_sample_layers: list[int] | None,
post_layer_norm: torch.nn.LayerNorm | None,
encoder_outputs: Union[torch.Tensor, list[torch.Tensor]],
feature_sample_layers: Optional[list[int]],
post_layer_norm: Optional[torch.nn.LayerNorm],
max_possible_layers: int,
) -> torch.Tensor:
"""Given the outputs a visual encoder module that may correspond to the
+8 -8
View File
@@ -20,14 +20,14 @@ import contextlib
import json
import os
from pathlib import Path
from typing import Any
from typing import Any, Dict, Optional, Type, Union
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: str | None = None,
model_override_args: dict | None = None,
revision: Optional[str] = None,
model_override_args: Optional[dict] = None,
**kwargs,
):
is_gguf = check_gguf_file(model)
@@ -83,8 +83,8 @@ def get_hf_config(
def get_diffusers_config(
model: str,
fastvideo_args: dict | None = None,
) -> dict[str, Any]:
fastvideo_args: Optional[dict] = 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: str | os.PathLike) -> bool:
def check_gguf_file(model: Union[str, os.PathLike]) -> bool:
"""Check if the file is a GGUF model."""
model = Path(model)
if not model.is_file():
@@ -6,8 +6,7 @@ import json
import os
import time
from abc import ABC, abstractmethod
from collections.abc import Generator, Iterable
from typing import Any, cast
from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
import torch
import torch.nn as nn
@@ -106,7 +105,7 @@ class TextEncoderLoader(ComponentLoader):
fall_back_to_pt: bool = True
"""Whether .pt weights can be used."""
allow_patterns_overrides: list[str] | None = None
allow_patterns_overrides: Optional[list[str]] = None
"""If defined, weights will load exclusively using these patterns."""
counter_before_loading_weights: float = 0.0
@@ -116,8 +115,8 @@ class TextEncoderLoader(ComponentLoader):
self,
model_name_or_path: str,
fall_back_to_pt: bool,
allow_patterns_overrides: list[str] | None,
) -> tuple[str, list[str], bool]:
allow_patterns_overrides: Optional[list[str]],
) -> Tuple[str, List[str], bool]:
"""Prepare weights for the model.
If the model is not local, it will be downloaded."""
@@ -139,7 +138,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 +161,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 +181,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="",
@@ -372,7 +371,6 @@ class TransformerLoader(ComponentLoader):
raise ValueError(
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
config.pop("_diffusers_version")
# Config from Diffusers supersedes fastvideo's model config
dit_config = fastvideo_args.dit_config
+11 -11
View File
@@ -7,9 +7,9 @@
import contextlib
import re
from collections import defaultdict
from collections.abc import Callable, Generator, Hashable
from itertools import chain
from typing import Any
from typing import (Any, Callable, DefaultDict, Dict, Generator, Hashable, List,
Optional, Tuple, Type)
import torch
from torch import nn
@@ -52,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.
@@ -87,12 +87,12 @@ def get_param_names_mapping(
# TODO(PY): add compile option
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,
cpu_offload: bool = False,
default_dtype: torch.dtype | None = torch.bfloat16,
default_dtype: Optional[torch.dtype] = torch.bfloat16,
) -> torch.nn.Module:
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
@@ -129,7 +129,7 @@ def shard_model(
*,
cpu_offload: bool,
reshard_after_forward: bool = True,
dp_mesh: DeviceMesh | None = None,
dp_mesh: Optional[DeviceMesh] = None,
) -> None:
"""
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API.
@@ -184,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: Callable[[str], tuple[str, Any, Any]] | None = None,
param_names_mapping: Optional[Callable[[str], tuple[str, Any, Any]]] = None,
) -> _IncompatibleKeys:
"""
Converting full state dict into a sharded state dict
@@ -212,7 +212,7 @@ def load_fsdp_model_from_full_model_state_dict(
meta_sharded_sd = model.state_dict()
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(
+17 -16
View File
@@ -8,8 +8,8 @@ import os
import tempfile
import time
from collections import defaultdict
from collections.abc import Generator
from pathlib import Path
from typing import Generator, List, Optional, Tuple, Union
import filelock
import huggingface_hub.constants
@@ -50,7 +50,8 @@ class DisabledTqdm(tqdm):
super().__init__(*args, **kwargs, disable=True)
def get_lock(model_name_or_path: str | Path, cache_dir: str | None = None):
def get_lock(model_name_or_path: Union[str, Path],
cache_dir: Optional[str] = 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)
@@ -76,10 +77,10 @@ def _shared_pointers(tensors):
def download_weights_from_hf(
model_name_or_path: str,
cache_dir: str | None,
allow_patterns: list[str],
revision: str | None = None,
ignore_patterns: str | list[str] | None = None,
cache_dir: Optional[str],
allow_patterns: List[str],
revision: Optional[str] = None,
ignore_patterns: Optional[Union[str, List[str]]] = None,
) -> str:
"""Download model weights from Hugging Face Hub.
@@ -135,8 +136,8 @@ def download_weights_from_hf(
def download_safetensors_index_file_from_hf(
model_name_or_path: str,
index_file: str,
cache_dir: str | None,
revision: str | None = None,
cache_dir: Optional[str],
revision: Optional[str] = None,
) -> None:
"""Download hf safetensors index file from Hugging Face Hub.
@@ -171,9 +172,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)
@@ -196,7 +197,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.
@@ -224,8 +225,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
@@ -242,8 +243,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
@@ -279,7 +280,7 @@ def default_weight_loader(param: torch.Tensor,
raise
def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> str | None:
def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> Optional[str]:
"""Remap the name of FP8 k/v_scale parameters.
This function handles the remapping of FP8 k/v_scale parameter names.
+13 -12
View File
@@ -1,9 +1,8 @@
# 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
from typing import Any, Callable, Tuple, Union
import torch
from torch.nn import Parameter
@@ -113,8 +112,9 @@ 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,8 +141,9 @@ 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)
@@ -229,7 +230,7 @@ class PerTensorScaleParameter(BasevLLMParameter):
self.qkv_idxs = {"q": 0, "k": 1, "v": 2}
super().__init__(**kwargs)
def _shard_id_as_int(self, shard_id: str | int) -> int:
def _shard_id_as_int(self, shard_id: Union[str, int]) -> int:
if isinstance(shard_id, int):
return shard_id
@@ -254,7 +255,7 @@ class PerTensorScaleParameter(BasevLLMParameter):
super().load_row_parallel_weight(*args, **kwargs)
def _load_into_shard_id(self, loaded_weight: torch.Tensor,
shard_id: str | int, **kwargs):
shard_id: Union[str, int], **kwargs):
"""
Slice the parameter data based on the shard id for
loading.
@@ -281,7 +282,7 @@ class PackedColumnParameter(_ColumnvLLMParameter):
for more details on the packed properties.
"""
def __init__(self, packed_factor: int | Fraction, packed_dim: int,
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
**kwargs):
self._packed_factor = packed_factor
self._packed_dim = packed_dim
@@ -296,7 +297,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,
@@ -314,7 +315,7 @@ class PackedvLLMParameter(ModelWeightParameter):
by accounting for packing and optionally, marlin tile size.
"""
def __init__(self, packed_factor: int | Fraction, packed_dim: int,
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
**kwargs):
self._packed_factor = packed_factor
self._packed_dim = packed_dim
@@ -403,7 +404,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
+30 -26
View File
@@ -8,10 +8,10 @@ import subprocess
import sys
import tempfile
from abc import ABC, abstractmethod
from collections.abc import Callable, Set
from dataclasses import dataclass, field
from functools import lru_cache
from typing import NoReturn, TypeVar, cast
from typing import (AbstractSet, Callable, Dict, List, NoReturn, Optional,
Tuple, Type, TypeVar, Union, cast)
import cloudpickle
from torch import nn
@@ -23,7 +23,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"HunyuanVideoTransformer3DModel":
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel")
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
"FluxTransformer2DModel": ("dits", "flux", "FluxTransformer2DModel"),
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
@@ -37,6 +38,7 @@ _TEXT_ENCODER_MODELS = {
"UMT5EncoderModel": ("encoders", "t5", "UMT5EncoderModel"),
"STEP1TextEncoder": ("encoders", "stepllm", "STEP1TextEncoder"),
"BertModel": ("encoders", "clip", "CLIPTextModel"),
"T5EncoderModel": ("encoders", "t5", "T5EncoderModel"),
}
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
@@ -48,13 +50,15 @@ _VAE_MODELS = {
"AutoencoderKLHunyuanVideo":
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo")
"AutoencoderKLStepvideo":
("vaes", "stepvideovae", "AutoencoderKLStepvideo"),
"AutoencoderKL": ("vaes", "image_vae", "AutoencoderKL"),
}
_SCHEDULERS = {
"FlowMatchEulerDiscreteScheduler":
("schedulers", "scheduling_flow_match_euler_discrete",
"FlowMatchDiscreteScheduler"),
"FlowMatchEulerDiscreteScheduler"),
"UniPCMultistepScheduler":
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
}
@@ -80,7 +84,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 +95,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 +106,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 +118,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 +163,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,
) -> type[nn.Module] | None:
) -> Optional[Type[nn.Module]]:
from fastvideo.v1.platforms import current_platform
current_platform.verify_model_arch(model_arch)
try:
@@ -182,7 +186,7 @@ def _try_load_model_cls(
def _try_inspect_model_cls(
model_arch: str,
model: _BaseRegisteredModel,
) -> _ModelInfo | None:
) -> Optional[_ModelInfo]:
try:
return model.inspect_model_cls()
except Exception:
@@ -194,15 +198,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) -> Set[str]:
def get_supported_archs(self) -> AbstractSet[str]:
return self.models.keys()
def register_model(
self,
model_arch: str,
model_cls: type[nn.Module] | str,
model_cls: Union[Type[nn.Module], str],
) -> None:
"""
Register an external model to be used in vLLM.
@@ -232,7 +236,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 +248,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) -> type[nn.Module] | None:
def _try_load_model_cls(self, model_arch: str) -> Optional[Type[nn.Module]]:
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) -> _ModelInfo | None:
def _try_inspect_model_cls(self, model_arch: str) -> Optional[_ModelInfo]:
if model_arch not in self.models:
return None
@@ -258,8 +262,8 @@ class _ModelRegistry:
def _normalize_archs(
self,
architectures: str | list[str],
) -> list[str]:
architectures: Union[str, List[str]],
) -> List[str]:
if isinstance(architectures, str):
architectures = [architectures]
if not architectures:
@@ -274,8 +278,8 @@ class _ModelRegistry:
def inspect_model_cls(
self,
architectures: str | list[str],
) -> tuple[_ModelInfo, str]:
architectures: Union[str, List[str]],
) -> Tuple[_ModelInfo, str]:
architectures = self._normalize_archs(architectures)
for arch in architectures:
@@ -287,8 +291,8 @@ class _ModelRegistry:
def resolve_model_cls(
self,
architectures: str | list[str],
) -> tuple[type[nn.Module], str]:
architectures: Union[str, List[str]],
) -> Tuple[Type[nn.Module], str]:
architectures = self._normalize_archs(architectures)
for arch in architectures:
+4 -3
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import Optional, Tuple, Union
import torch
from diffusers.utils import BaseOutput
@@ -31,15 +32,15 @@ class BaseScheduler(ABC):
@abstractmethod
def scale_model_input(self,
sample: torch.Tensor,
timestep: int | None = None) -> torch.Tensor:
timestep: Optional[int] = None) -> torch.Tensor:
pass
@abstractmethod
def step(
self,
model_output: torch.Tensor,
timestep: int | torch.Tensor,
timestep: Union[int, torch.Tensor],
sample: torch.Tensor,
return_dict: bool = True,
) -> BaseOutput | tuple:
) -> Union[BaseOutput, Tuple]:
pass
@@ -1,3 +1,4 @@
# type: ignore
# SPDX-License-Identifier: Apache-2.0
# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved.
@@ -19,13 +20,16 @@
#
# ==============================================================================
import math
from dataclasses import dataclass
from typing import Any
from typing import List, Optional, Tuple, Union
import numpy as np
import scipy
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
from diffusers.utils import BaseOutput, is_scipy_available, logging
from fastvideo.v1.models.schedulers.base import BaseScheduler
@@ -33,7 +37,7 @@ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@dataclass
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
class FlowMatchEulerDiscreteSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
@@ -46,7 +50,8 @@ class FlowMatchDiscreteSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
BaseScheduler):
"""
Euler scheduler.
@@ -56,16 +61,37 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps to train the model.
timestep_spacing (`str`, defaults to `"linspace"`):
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
shift (`float`, defaults to 1.0):
The shift value for the timestep schedule.
reverse (`bool`, defaults to `True`):
Whether to reverse the timestep schedule.
use_dynamic_shifting (`bool`, defaults to False):
Whether to apply timestep shifting on-the-fly based on the image resolution.
base_shift (`float`, defaults to 0.5):
Value to stabilize image generation. Increasing `base_shift` reduces variation and image is more consistent
with desired output.
max_shift (`float`, defaults to 1.15):
Value change allowed to latent vectors. Increasing `max_shift` encourages more variation and image may be
more exaggerated or stylized.
base_image_seq_len (`int`, defaults to 256):
The base image sequence length.
max_image_seq_len (`int`, defaults to 4096):
The maximum image sequence length.
invert_sigmas (`bool`, defaults to False):
Whether to invert the sigmas.
shift_terminal (`float`, defaults to None):
The end value of the shifted timestep schedule.
use_karras_sigmas (`bool`, defaults to False):
Whether to use Karras sigmas for step sizes in the noise schedule during sampling.
use_exponential_sigmas (`bool`, defaults to False):
Whether to use exponential sigmas for step sizes in the noise schedule during sampling.
use_beta_sigmas (`bool`, defaults to False):
Whether to use beta sigmas for step sizes in the noise schedule during sampling.
time_shift_type (`str`, defaults to "exponential"):
The type of dynamic resolution-dependent timestep shifting to apply. Either "exponential" or "linear".
stochastic_sampling (`bool`, defaults to False):
Whether to use stochastic sampling.
"""
_compatibles: list[Any] = []
_compatibles = []
order = 1
@register_to_config
@@ -73,31 +99,62 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self,
num_train_timesteps: int = 1000,
shift: float = 1.0,
reverse: bool = True,
solver: str = "euler",
n_tokens: int | None = None,
**kwargs,
use_dynamic_shifting: bool = False,
base_shift: Optional[float] = 0.5,
max_shift: Optional[float] = 1.15,
base_image_seq_len: Optional[int] = 256,
max_image_seq_len: Optional[int] = 4096,
invert_sigmas: bool = False,
shift_terminal: Optional[float] = None,
use_karras_sigmas: Optional[bool] = False,
use_exponential_sigmas: Optional[bool] = False,
use_beta_sigmas: Optional[bool] = False,
time_shift_type: str = "exponential",
stochastic_sampling: bool = False,
):
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
if not reverse:
sigmas = sigmas.flip(0)
self.sigmas = sigmas
# the value fed to model
self.timesteps = (sigmas[:-1] *
num_train_timesteps).to(dtype=torch.float32)
self._step_index: int | None = None
self._begin_index = 0
self.supported_solver = ["euler"]
if solver not in self.supported_solver:
if self.config.use_beta_sigmas and not is_scipy_available():
raise ImportError(
"Make sure to install scipy if you want to use beta sigmas.")
if sum([
self.config.use_beta_sigmas, self.config.use_exponential_sigmas,
self.config.use_karras_sigmas
]) > 1:
raise ValueError(
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
"Only one of `config.use_beta_sigmas`, `config.use_exponential_sigmas`, `config.use_karras_sigmas` can be used."
)
if time_shift_type not in {"exponential", "linear"}:
raise ValueError(
"`time_shift_type` must either be 'exponential' or 'linear'.")
BaseScheduler.__init__(self)
timesteps = np.linspace(1,
num_train_timesteps,
num_train_timesteps,
dtype=np.float32)[::-1].copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
sigmas = timesteps / num_train_timesteps
if not use_dynamic_shifting:
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
self.timesteps = sigmas * num_train_timesteps
self._step_index = None
self._begin_index = None
self._shift = shift
self.sigmas = sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
@property
def shift(self):
"""
The value used for shifting.
"""
return self._shift
@property
def step_index(self):
@@ -124,34 +181,197 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
"""
self._begin_index = begin_index
def set_shift(self, shift: float):
self._shift = shift
def scale_noise(
self,
sample: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
noise: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
"""
Forward process in flow-matching
Args:
sample (`torch.FloatTensor`):
The input sample.
timestep (`int`, *optional*):
The current timestep in the diffusion chain.
Returns:
`torch.FloatTensor`:
A scaled input sample.
"""
# Make sure sigmas and timesteps have the same device and dtype as original_samples
sigmas = self.sigmas.to(device=sample.device, dtype=sample.dtype)
if sample.device.type == "mps" and torch.is_floating_point(timestep):
# mps does not support float64
schedule_timesteps = self.timesteps.to(sample.device,
dtype=torch.float32)
timestep = timestep.to(sample.device, dtype=torch.float32)
else:
schedule_timesteps = self.timesteps.to(sample.device)
timestep = timestep.to(sample.device)
# self.begin_index is None when scheduler is used for training, or pipeline does not implement set_begin_index
if self.begin_index is None:
step_indices = [
self.index_for_timestep(t, schedule_timesteps) for t in timestep
]
elif self.step_index is not None:
# add_noise is called after first denoising step (for inpainting)
step_indices = [self.step_index] * timestep.shape[0]
else:
# add noise is called before first denoising step to create initial latent(img2img)
step_indices = [self.begin_index] * timestep.shape[0]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < len(sample.shape):
sigma = sigma.unsqueeze(-1)
sample = sigma * noise + (1.0 - sigma) * sample
return sample
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
if self.config.time_shift_type == "exponential":
return self._time_shift_exponential(mu, sigma, t)
elif self.config.time_shift_type == "linear":
return self._time_shift_linear(mu, sigma, t)
def stretch_shift_to_terminal(self, t: torch.Tensor) -> torch.Tensor:
r"""
Stretches and shifts the timestep schedule to ensure it terminates at the configured `shift_terminal` config
value.
Reference:
https://github.com/Lightricks/LTX-Video/blob/a01a171f8fe3d99dce2728d60a73fecf4d4238ae/ltx_video/schedulers/rf.py#L51
Args:
t (`torch.Tensor`):
A tensor of timesteps to be stretched and shifted.
Returns:
`torch.Tensor`:
A tensor of adjusted timesteps such that the final value equals `self.config.shift_terminal`.
"""
one_minus_z = 1 - t
scale_factor = one_minus_z[-1] / (1 - self.config.shift_terminal)
stretched_t = 1 - (one_minus_z / scale_factor)
return stretched_t
def set_timesteps(
self,
num_inference_steps: int,
device: str | torch.device = None,
n_tokens: int = 0,
num_inference_steps: Optional[int] = None,
device: Union[str, torch.device] = None,
sigmas: Optional[List[float]] = None,
mu: Optional[float] = None,
timesteps: Optional[List[float]] = None,
**kwargs,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`):
num_inference_steps (`int`, *optional*):
The number of diffusion steps used when generating samples with a pre-trained model.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
sigmas (`List[float]`, *optional*):
Custom values for sigmas to be used for each diffusion step. If `None`, the sigmas are computed
automatically.
mu (`float`, *optional*):
Determines the amount of shifting applied to sigmas when performing resolution-dependent timestep
shifting.
timesteps (`List[float]`, *optional*):
Custom values for timesteps to be used for each diffusion step. If `None`, the timesteps are computed
automatically.
"""
if self.config.use_dynamic_shifting and mu is None:
raise ValueError(
"`mu` must be passed when `use_dynamic_shifting` is set to be `True`"
)
if sigmas is not None and timesteps is not None and len(sigmas) != len(
timesteps):
raise ValueError(
"`sigmas` and `timesteps` should have the same length")
if num_inference_steps is not None:
if (sigmas is not None and len(sigmas) != num_inference_steps) or (
timesteps is not None
and len(timesteps) != num_inference_steps):
raise ValueError(
"`sigmas` and `timesteps` should have the same length as num_inference_steps, if `num_inference_steps` is provided"
)
else:
num_inference_steps = len(sigmas) if sigmas is not None else len(
timesteps)
self.num_inference_steps = num_inference_steps
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
sigmas = self.sd3_time_shift(sigmas)
# 1. Prepare default sigmas
is_timesteps_provided = timesteps is not None
if not self.config.reverse:
sigmas = 1 - sigmas
if is_timesteps_provided:
timesteps = np.array(timesteps).astype(np.float32)
if sigmas is None:
if timesteps is None:
timesteps = np.linspace(self._sigma_to_t(self.sigma_max),
self._sigma_to_t(self.sigma_min),
num_inference_steps)
sigmas = timesteps / self.config.num_train_timesteps
else:
sigmas = np.array(sigmas).astype(np.float32)
num_inference_steps = len(sigmas)
# 2. Perform timestep shifting. Either no shifting is applied, or resolution-dependent shifting of
# "exponential" or "linear" type is applied
if self.config.use_dynamic_shifting:
sigmas = self.time_shift(mu, 1.0, sigmas)
else:
sigmas = self.shift * sigmas / (1 + (self.shift - 1) * sigmas)
# 3. If required, stretch the sigmas schedule to terminate at the configured `shift_terminal` value
if self.config.shift_terminal:
sigmas = self.stretch_shift_to_terminal(sigmas)
# 4. If required, convert sigmas to one of karras, exponential, or beta sigma schedules
if self.config.use_karras_sigmas:
sigmas = self._convert_to_karras(
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
elif self.config.use_exponential_sigmas:
sigmas = self._convert_to_exponential(
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
elif self.config.use_beta_sigmas:
sigmas = self._convert_to_beta(
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
# 5. Convert sigmas and timesteps to tensors and move to specified device
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32, device=device)
if not is_timesteps_provided:
timesteps = sigmas * self.config.num_train_timesteps
else:
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32,
device=device)
# 6. Append the terminal sigma value.
# If a model requires inverted sigma schedule for denoising but timesteps without inversion, the
# `invert_sigmas` flag can be set to `True`. This case is only required in Mochi
if self.config.invert_sigmas:
sigmas = 1.0 - sigmas
timesteps = sigmas * self.config.num_train_timesteps
sigmas = torch.cat([sigmas, torch.ones(1, device=sigmas.device)])
else:
sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
self.timesteps = timesteps
self.sigmas = sigmas
if not getattr(self.config, "timesteps_scale", True):
self.timesteps = sigmas[:-1] # for stepvideo
@@ -160,8 +380,14 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
dtype=torch.float32, device=device)
# Reset step index
self._step_index = None
self._begin_index = None
def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
def scale_model_input(self,
sample: torch.Tensor,
timestep: Optional[int] = None) -> torch.Tensor:
return sample
def index_for_timestep(self, timestep, schedule_timesteps=None):
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
@@ -173,9 +399,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
pos = 1 if len(indices) > 1 else 0
idx: int = indices[pos].item()
return idx
return indices[pos].item()
def set_shift(self, shift: float) -> None:
self.config.shift = shift
@@ -191,22 +415,19 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
else:
self._step_index = self._begin_index
def scale_model_input(self,
sample: torch.Tensor,
timestep: int | None = None) -> torch.Tensor:
return sample
def sd3_time_shift(self, t: torch.Tensor):
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
def step(
self,
model_output: torch.FloatTensor,
timestep: float | torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
s_churn: float = 0.0,
s_tmin: float = 0.0,
s_tmax: float = float("inf"),
s_noise: float = 1.0,
generator: Optional[torch.Generator] = None,
per_token_timesteps: Optional[torch.Tensor] = None,
return_dict: bool = True,
**kwargs,
) -> FlowMatchDiscreteSchedulerOutput | tuple:
) -> Union[FlowMatchEulerDiscreteSchedulerOutput, 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).
@@ -218,24 +439,30 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
The current discrete timestep in the diffusion chain.
sample (`torch.FloatTensor`):
A current instance of a sample created by the diffusion process.
s_churn (`float`):
s_tmin (`float`):
s_tmax (`float`):
s_noise (`float`, defaults to 1.0):
Scaling factor for noise added to the sample.
generator (`torch.Generator`, *optional*):
A random number generator.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
per_token_timesteps (`torch.Tensor`, *optional*):
The timesteps for each token in the sample.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
tuple.
Whether or not to return a
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] or tuple.
Returns:
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
returned, otherwise a tuple is returned where the first element is the sample tensor.
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`,
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] is 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"
" `FlowMatchEulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."), )
if self.step_index is None:
@@ -244,24 +471,132 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
# Upcast to avoid precision issues when computing prev_sample
sample = sample.to(torch.float32)
assert self.step_index is not None
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
if per_token_timesteps is not None:
per_token_sigmas = per_token_timesteps / self.config.num_train_timesteps
if self.config.solver == "euler":
prev_sample = sample + model_output.to(torch.float32) * dt
sigmas = self.sigmas[:, None, None]
lower_mask = sigmas < per_token_sigmas[None] - 1e-6
lower_sigmas = lower_mask * sigmas
lower_sigmas, _ = lower_sigmas.max(dim=0)
current_sigma = per_token_sigmas[..., None]
next_sigma = lower_sigmas[..., None]
dt = current_sigma - next_sigma
else:
raise ValueError(
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
)
sigma_idx = self.step_index
sigma = self.sigmas[sigma_idx]
sigma_next = self.sigmas[sigma_idx + 1]
current_sigma = sigma
next_sigma = sigma_next
dt = sigma_next - sigma
if self.config.stochastic_sampling:
x0 = sample - current_sigma * model_output
noise = torch.randn_like(sample)
prev_sample = (1.0 - next_sigma) * x0 + next_sigma * noise
else:
prev_sample = sample + dt * model_output
# upon completion increase step index by one
assert self._step_index is not None
self._step_index += 1
if per_token_timesteps is None:
# Cast sample back to model compatible dtype
prev_sample = prev_sample.to(model_output.dtype)
if not return_dict:
return (prev_sample, )
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
return FlowMatchEulerDiscreteSchedulerOutput(prev_sample=prev_sample)
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_karras
def _convert_to_karras(self, in_sigmas: torch.Tensor,
num_inference_steps) -> torch.Tensor:
"""Constructs the noise schedule of Karras et al. (2022)."""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
rho = 7.0 # 7.0 is the value used in the paper
ramp = np.linspace(0, 1, num_inference_steps)
min_inv_rho = sigma_min**(1 / rho)
max_inv_rho = sigma_max**(1 / rho)
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho))**rho
return sigmas
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_exponential
def _convert_to_exponential(self, in_sigmas: torch.Tensor,
num_inference_steps: int) -> torch.Tensor:
"""Constructs an exponential noise schedule."""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
sigmas = np.exp(
np.linspace(math.log(sigma_max), math.log(sigma_min),
num_inference_steps))
return sigmas
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_beta
def _convert_to_beta(self,
in_sigmas: torch.Tensor,
num_inference_steps: int,
alpha: float = 0.6,
beta: float = 0.6) -> torch.Tensor:
"""From "Beta Sampling is All You Need" [arXiv:2407.12173] (Lee et. al, 2024)"""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
sigmas = np.array([
sigma_min + (ppf * (sigma_max - sigma_min)) for ppf in [
scipy.stats.beta.ppf(timestep, alpha, beta)
for timestep in 1 - np.linspace(0, 1, num_inference_steps)
]
])
return sigmas
def _time_shift_exponential(self, mu, sigma, t):
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma)
def _time_shift_linear(self, mu, sigma, t):
return mu / (mu + (1 / t - 1)**sigma)
def __len__(self):
return self.config.num_train_timesteps
@@ -23,6 +23,7 @@
# ==============================================================================
import math
from typing import List, Optional, Tuple, Union
import numpy as np
import torch
@@ -202,7 +203,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
beta_start: float = 0.0001,
beta_end: float = 0.02,
beta_schedule: str = "linear",
trained_betas: np.ndarray | list[float] | None = None,
trained_betas: Optional[Union[np.ndarray, List[float]]] = None,
solver_order: int = 2,
prediction_type: str = "epsilon",
thresholding: bool = False,
@@ -211,16 +212,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: 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,
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,
timestep_spacing: str = "linspace",
steps_offset: int = 0,
final_sigmas_type: str | None = "zero", # "zero", "sigma_min"
final_sigmas_type: Optional[str] = "zero", # "zero", "sigma_min"
rescale_betas_zero_snr: bool = False,
):
if self.config.use_beta_sigmas and not is_scipy_available():
@@ -282,20 +283,21 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self.predict_x0 = predict_x0
# setable values
self.num_inference_steps: int | None = None
self.num_inference_steps: Optional[int] = 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[int | torch.Tensor] = [None] * solver_order
self.timestep_list: List[Union[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: int | None = None
self._begin_index: int | None = None
self._step_index: Optional[int] = None
self._begin_index: Optional[int] = None
self.sigmas = self.sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
@@ -331,7 +333,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def set_timesteps(self,
num_inference_steps: int,
device: str | torch.device = None):
device: Union[str, torch.device] = None):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
@@ -535,7 +537,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
@@ -706,7 +708,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
model_output: torch.Tensor,
*args,
sample: torch.Tensor = None,
order: int | None = None,
order: Optional[int] = None,
**kwargs,
) -> torch.Tensor:
"""
@@ -806,7 +808,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
R_tensor: torch.Tensor = torch.stack(R)
b = torch.tensor(b, device=device)
D1s_tensor: torch.Tensor | None = None
D1s_tensor: Optional[torch.Tensor] = None
if len(D1s) > 0:
D1s_tensor = torch.stack(D1s, dim=1) # (B, K)
# for order 2, we use a simplified version
@@ -840,9 +842,9 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self,
this_model_output: torch.Tensor,
*args,
last_sample: torch.Tensor | None = None,
this_sample: torch.Tensor | None = None,
order: int | None = None,
last_sample: Optional[torch.Tensor] = None,
this_sample: Optional[torch.Tensor] = None,
order: Optional[int] = None,
**kwargs,
) -> torch.Tensor:
"""
@@ -948,7 +950,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
R = torch.stack(R)
b = torch.tensor(b, device=device)
D1s_tensor: torch.Tensor | None = torch.stack(
D1s_tensor: Optional[torch.Tensor] = torch.stack(
D1s, dim=1) if len(D1s) > 0 else None
# for order 1, we use a simplified version
@@ -1014,10 +1016,10 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def step(
self,
model_output: torch.Tensor,
timestep: int | torch.Tensor,
timestep: Union[int, torch.Tensor],
sample: torch.Tensor,
return_dict: bool = True,
) -> SchedulerOutput | tuple:
) -> Union[SchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with
the multistep UniPC.
+5 -5
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/utils.py
"""Utils for model executor."""
from typing import Any
from typing import Any, Dict, List, Optional
import torch
@@ -58,7 +58,7 @@ def set_random_seed(seed: int) -> None:
def set_weight_attrs(
weight: torch.Tensor,
weight_attrs: dict[str, Any] | None,
weight_attrs: Optional[Dict[str, Any]],
):
"""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: torch.Tensor | None = None,
scale: torch.Tensor | None = None) -> torch.Tensor:
shift: Optional[torch.Tensor] = None,
scale: Optional[torch.Tensor] = None) -> torch.Tensor:
"""modulate by shift and scale
Args:
+17 -20
View File
@@ -1,9 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from collections.abc import Iterator
from math import prod
from typing import Optional, cast
from typing import Iterator, Optional, Tuple, Union, cast
import numpy as np
import torch
@@ -40,9 +39,6 @@ class ParallelTiledVAE(ABC):
self.use_temporal_tiling = config.use_temporal_tiling
self.use_parallel_tiling = config.use_parallel_tiling
def to(self, device) -> 'ParallelTiledVAE':
return self
@property
def temporal_compression_ratio(self) -> int:
return cast(int, self.config.temporal_compression_ratio)
@@ -52,8 +48,8 @@ class ParallelTiledVAE(ABC):
return cast(int, self.config.spatial_compression_ratio)
@property
def scaling_factor(self) -> float | torch.Tensor:
return cast(float | torch.Tensor, self.config.scaling_factor)
def scaling_factor(self) -> Union[float, torch.tensor]:
return cast(Union[float, torch.tensor], self.config.scaling_factor)
@abstractmethod
def _encode(self, *args, **kwargs) -> torch.Tensor:
@@ -161,7 +157,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
@@ -413,16 +409,16 @@ class ParallelTiledVAE(ABC):
def enable_tiling(
self,
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,
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,
) -> None:
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
@@ -486,7 +482,8 @@ class DiagonalGaussianDistribution:
device=self.parameters.device,
dtype=self.parameters.dtype)
def sample(self, generator: torch.Generator | None = None) -> torch.Tensor:
def sample(self,
generator: Optional[torch.Generator] = 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 +514,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)
+23 -25
View File
@@ -15,6 +15,8 @@
# 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
@@ -30,7 +32,7 @@ def prepare_causal_attention_mask(
height_width: int,
dtype: torch.dtype,
device: torch.device,
batch_size: int | None = None) -> torch.Tensor:
batch_size: Optional[int] = 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")
@@ -70,7 +72,7 @@ class HunyuanVAEAttention(nn.Module):
def forward(self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None) -> torch.Tensor:
attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
residual = hidden_states
batch_size, sequence_length, _ = hidden_states.shape
@@ -119,10 +121,10 @@ class HunyuanVideoCausalConv3d(nn.Module):
self,
in_channels: int,
out_channels: int,
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,
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,
bias: bool = True,
pad_mode: str = "replicate",
) -> None:
@@ -161,11 +163,11 @@ class HunyuanVideoUpsampleCausal3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int | None = None,
out_channels: Optional[int] = 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__()
@@ -211,7 +213,7 @@ class HunyuanVideoDownsampleCausal3D(nn.Module):
def __init__(
self,
channels: int,
out_channels: int | None = None,
out_channels: Optional[int] = None,
padding: int = 1,
kernel_size: int = 3,
bias: bool = True,
@@ -237,7 +239,7 @@ class HunyuanVideoResnetBlockCausal3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int | None = None,
out_channels: Optional[int] = None,
dropout: float = 0.0,
groups: int = 32,
eps: float = 1e-6,
@@ -311,7 +313,7 @@ class HunyuanVideoMidBlock3D(nn.Module):
non_linearity=resnet_act_fn,
)
]
attentions: list[HunyuanVAEAttention | None] = []
attentions: list[Optional[HunyuanVAEAttention]] = []
for _ in range(num_layers):
if self.add_attention:
@@ -347,9 +349,7 @@ 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:],
strict=False):
for attn, resnet in zip(self.attentions, self.resnets[1:]):
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,9 +371,7 @@ class HunyuanVideoMidBlock3D(nn.Module):
else:
hidden_states = self.resnets[0](hidden_states)
for attn, resnet in zip(self.attentions,
self.resnets[1:],
strict=False):
for attn, resnet in zip(self.attentions, self.resnets[1:]):
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,
@@ -406,7 +404,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__()
@@ -468,7 +466,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 = []
@@ -527,13 +525,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",
@@ -548,7 +546,7 @@ class HunyuanVideoEncoder3D(nn.Module):
block_out_channels[0],
kernel_size=3,
stride=1)
self.mid_block: HunyuanVideoMidBlock3D | None = None
self.mid_block: Optional[HunyuanVideoMidBlock3D] = None
self.down_blocks = nn.ModuleList([])
output_channel = block_out_channels[0]
@@ -651,13 +649,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",
@@ -834,7 +832,7 @@ class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
self,
sample: torch.Tensor,
sample_posterior: bool = False,
generator: torch.Generator | None = None,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
r"""
Args:
+556
View File
@@ -0,0 +1,556 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from diffusers
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Dict, Optional, Union
import torch
import torch.nn as nn
from diffusers.models.attention_processor import (ADDED_KV_ATTENTION_PROCESSORS,
CROSS_ATTENTION_PROCESSORS,
Attention, AttentionProcessor,
AttnAddedKVProcessor,
AttnProcessor,
FusedAttnProcessor2_0)
from diffusers.models.autoencoders.vae import Decoder, Encoder
from fastvideo.v1.configs.models.vaes import ImageVAEConfig
from fastvideo.v1.models.vaes.common import DiagonalGaussianDistribution
class AutoencoderKL(nn.Module):
r"""
A VAE model with KL loss for encoding images into latents and decoding latent representations into images.
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
for all models (such as downloading or saving).
Parameters:
in_channels (int, *optional*, defaults to 3): Number of channels in the input image.
out_channels (int, *optional*, defaults to 3): Number of channels in the output.
down_block_types (`Tuple[str]`, *optional*, defaults to `("DownEncoderBlock2D",)`):
Tuple of downsample block types.
up_block_types (`Tuple[str]`, *optional*, defaults to `("UpDecoderBlock2D",)`):
Tuple of upsample block types.
block_out_channels (`Tuple[int]`, *optional*, defaults to `(64,)`):
Tuple of block output channels.
act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use.
latent_channels (`int`, *optional*, defaults to 4): Number of channels in the latent space.
sample_size (`int`, *optional*, defaults to `32`): Sample input size.
scaling_factor (`float`, *optional*, defaults to 0.18215):
The component-wise standard deviation of the trained latent space computed using the first batch of the
training set. This is used to scale the latent space to have unit variance when training the diffusion
model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the
diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1
/ scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image
Synthesis with Latent Diffusion Models](https://arxiv.org/abs/2112.10752) paper.
force_upcast (`bool`, *optional*, default to `True`):
If enabled it will force the VAE to run in float32 for high image resolution pipelines, such as SD-XL. VAE
can be fine-tuned / trained to a lower range without losing too much precision in which case
`force_upcast` can be set to `False` - see: https://huggingface.co/madebyollin/sdxl-vae-fp16-fix
mid_block_add_attention (`bool`, *optional*, default to `True`):
If enabled, the mid_block of the Encoder and Decoder will have attention blocks. If set to false, the
mid_block will only have resnet blocks
"""
_supports_gradient_checkpointing = True
_no_split_modules = ["BasicTransformerBlock", "ResnetBlock2D"]
def __init__(
self,
config: ImageVAEConfig,
):
nn.Module.__init__(self)
self.shift_factor = config.shift_factor
self.scaling_factor = config.scaling_factor
if config.load_encoder:
# pass init params to Encoder
self.encoder = Encoder(
in_channels=config.in_channels,
out_channels=config.latent_channels,
down_block_types=config.down_block_types,
block_out_channels=config.block_out_channels,
layers_per_block=config.layers_per_block,
act_fn=config.act_fn,
norm_num_groups=config.norm_num_groups,
double_z=True,
mid_block_add_attention=config.mid_block_add_attention,
)
self.quant_conv = nn.Conv2d(2 * config.latent_channels, 2 *
config.latent_channels,
1) if config.use_quant_conv else None
if config.load_decoder:
# pass init params to Decoder
self.decoder = Decoder(
in_channels=config.latent_channels,
out_channels=config.out_channels,
up_block_types=config.up_block_types,
block_out_channels=config.block_out_channels,
layers_per_block=config.layers_per_block,
norm_num_groups=config.norm_num_groups,
act_fn=config.act_fn,
mid_block_add_attention=config.mid_block_add_attention,
)
self.post_quant_conv = nn.Conv2d(
config.latent_channels, config.latent_channels,
1) if config.use_post_quant_conv else None
self.use_slicing = False
self.use_tiling = False
# only relevant if vae tiling is enabled
self.tile_sample_min_size = config.sample_size
sample_size = (config.sample_size[0] if isinstance(
config.sample_size, (list, tuple)) else config.sample_size)
self.tile_latent_min_size = int(
sample_size / (2**(len(config.block_out_channels) - 1)))
self.tile_overlap_factor = 0.25
def enable_tiling(self, use_tiling: bool = True):
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
processing larger images.
"""
self.use_tiling = use_tiling
def disable_tiling(self):
r"""
Disable tiled VAE decoding. If `enable_tiling` was previously enabled, this method will go back to computing
decoding in one step.
"""
self.enable_tiling(False)
def enable_slicing(self):
r"""
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
"""
self.use_slicing = True
def disable_slicing(self):
r"""
Disable sliced VAE decoding. If `enable_slicing` was previously enabled, this method will go back to computing
decoding in one step.
"""
self.use_slicing = False
@property
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors
def attn_processors(self) -> Dict[str, AttentionProcessor]:
r"""
Returns:
`dict` of attention processors: A dictionary containing all attention processors used in the model with
indexed by its weight name.
"""
# set recursively
processors: dict[str, AttentionProcessor] = {}
def fn_recursive_add_processors(name: str, module: torch.nn.Module,
processors: Dict[str,
AttentionProcessor]):
if hasattr(module, "get_processor"):
processors[f"{name}.processor"] = module.get_processor()
for sub_name, child in module.named_children():
fn_recursive_add_processors(f"{name}.{sub_name}", child,
processors)
return processors
for name, module in self.named_children():
fn_recursive_add_processors(name, module, processors)
return processors
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor
def set_attn_processor(self, processor: Union[AttentionProcessor,
Dict[str,
AttentionProcessor]]):
r"""
Sets the attention processor to use to compute attention.
Parameters:
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
The instantiated processor class or a dictionary of processor classes that will be set as the processor
for **all** `Attention` layers.
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
processor. This is strongly recommended when setting trainable attention processors.
"""
count = len(self.attn_processors.keys())
if isinstance(processor, dict) and len(processor) != count:
raise ValueError(
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
)
def fn_recursive_attn_processor(name: str, module: torch.nn.Module,
processor):
if hasattr(module, "set_processor"):
if not isinstance(processor, dict):
module.set_processor(processor)
else:
module.set_processor(processor.pop(f"{name}.processor"))
for sub_name, child in module.named_children():
fn_recursive_attn_processor(f"{name}.{sub_name}", child,
processor)
for name, module in self.named_children():
fn_recursive_attn_processor(name, module, processor)
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor
def set_default_attn_processor(self):
"""
Disables custom attention processors and sets the default attention implementation.
"""
if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS
for proc in self.attn_processors.values()):
processor = AttnAddedKVProcessor()
elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS
for proc in self.attn_processors.values()):
processor = AttnProcessor()
else:
raise ValueError(
f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}"
)
self.set_attn_processor(processor)
def _encode(self, x: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, height, width = x.shape
if self.use_tiling and (width > self.tile_sample_min_size
or height > self.tile_sample_min_size):
return self._tiled_encode(x)
enc = self.encoder(x)
if self.quant_conv is not None:
enc = self.quant_conv(enc)
return enc
def encode(self, x: torch.Tensor) -> torch.Tensor:
"""
Encode a batch of images into latents.
Args:
x (`torch.Tensor`): Input batch of images.
return_dict (`bool`, *optional*, defaults to `True`):
Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
Returns:
The latent representations of the encoded images. If `return_dict` is True, a
[`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned.
"""
if self.use_slicing and x.shape[0] > 1:
encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)]
enc = torch.cat(encoded_slices)
else:
enc = self._encode(x)
enc = DiagonalGaussianDistribution(enc)
return enc
def _decode(self, z: torch.Tensor) -> torch.Tensor:
if self.use_tiling and (z.shape[-1] > self.tile_latent_min_size
or z.shape[-2] > self.tile_latent_min_size):
return self.tiled_decode(z)
if self.post_quant_conv is not None:
z = self.post_quant_conv(z)
dec = self.decoder(z)
return dec
def decode(self, z: torch.Tensor) -> torch.Tensor:
"""
Decode a batch of images.
Args:
z (`torch.Tensor`): Input batch of latent vectors.
return_dict (`bool`, *optional*, defaults to `True`):
Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
Returns:
[`~models.vae.DecoderOutput`] or `tuple`:
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
returned.
"""
if self.use_slicing and z.shape[0] > 1:
decoded_slices = [
self._decode(z_slice).sample for z_slice in z.split(1)
]
decoded = torch.cat(decoded_slices)
else:
decoded = self._decode(z)
return decoded
def blend_v(self, a: torch.Tensor, b: torch.Tensor,
blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[2], b.shape[2], blend_extent)
for y in range(blend_extent):
b[:, :, y, :] = a[:, :, -blend_extent + y, :] * (
1 - y / blend_extent) + b[:, :, y, :] * (y / blend_extent)
return b
def blend_h(self, a: torch.Tensor, b: torch.Tensor,
blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[3], b.shape[3], blend_extent)
for x in range(blend_extent):
b[:, :, :, x] = a[:, :, :, -blend_extent + x] * (
1 - x / blend_extent) + b[:, :, :, x] * (x / blend_extent)
return b
def _tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
r"""Encode a batch of images using a tiled encoder.
When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is
different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
output, but they should be much less noticeable.
Args:
x (`torch.Tensor`): Input batch of images.
Returns:
`torch.Tensor`:
The latent representation of the encoded videos.
"""
overlap_size = int(self.tile_sample_min_size *
(1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
row_limit = self.tile_latent_min_size - blend_extent
# Split the image into 512x512 tiles and encode them separately.
rows = []
for i in range(0, x.shape[2], overlap_size):
row = []
for j in range(0, x.shape[3], overlap_size):
tile = x[:, :, i:i + self.tile_sample_min_size,
j:j + self.tile_sample_min_size]
tile = self.encoder(tile)
if self.quant_conv:
tile = self.quant_conv(tile)
row.append(tile)
rows.append(row)
result_rows = []
for i, row in enumerate(rows):
result_row = []
for j, tile in enumerate(row):
# blend the above tile and the left tile
# to the current tile and add the current tile to the result row
if i > 0:
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :row_limit, :row_limit])
result_rows.append(torch.cat(result_row, dim=3))
enc = torch.cat(result_rows, dim=2)
return enc
def tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
r"""Encode a batch of images using a tiled encoder.
When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is
different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
output, but they should be much less noticeable.
Args:
x (`torch.Tensor`): Input batch of images.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
Returns:
[`~models.autoencoder_kl.AutoencoderKLOutput`] or `tuple`:
If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain
`tuple` is returned.
"""
overlap_size = int(self.tile_sample_min_size *
(1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
row_limit = self.tile_latent_min_size - blend_extent
# Split the image into 512x512 tiles and encode them separately.
rows = []
for i in range(0, x.shape[2], overlap_size):
row = []
for j in range(0, x.shape[3], overlap_size):
tile = x[:, :, i:i + self.tile_sample_min_size,
j:j + self.tile_sample_min_size]
tile = self.encoder(tile)
if self.quant_conv:
tile = self.quant_conv(tile)
row.append(tile)
rows.append(row)
result_rows = []
for i, row in enumerate(rows):
result_row = []
for j, tile in enumerate(row):
# blend the above tile and the left tile
# to the current tile and add the current tile to the result row
if i > 0:
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :row_limit, :row_limit])
result_rows.append(torch.cat(result_row, dim=3))
moments = torch.cat(result_rows, dim=2)
enc = DiagonalGaussianDistribution(moments)
return enc
def tiled_decode(self, z: torch.Tensor) -> torch.Tensor:
r"""
Decode a batch of images using a tiled decoder.
Args:
z (`torch.Tensor`): Input batch of latent vectors.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
Returns:
[`~models.vae.DecoderOutput`] or `tuple`:
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
returned.
"""
overlap_size = int(self.tile_latent_min_size *
(1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
row_limit = self.tile_sample_min_size - blend_extent
# Split z into overlapping 64x64 tiles and decode them separately.
# The tiles have an overlap to avoid seams between tiles.
rows = []
for i in range(0, z.shape[2], overlap_size):
row = []
for j in range(0, z.shape[3], overlap_size):
tile = z[:, :, i:i + self.tile_latent_min_size,
j:j + self.tile_latent_min_size]
if self.post_quant_conv:
tile = self.post_quant_conv(tile)
decoded = self.decoder(tile)
row.append(decoded)
rows.append(row)
result_rows = []
for i, row in enumerate(rows):
result_row = []
for j, tile in enumerate(row):
# blend the above tile and the left tile
# to the current tile and add the current tile to the result row
if i > 0:
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :row_limit, :row_limit])
result_rows.append(torch.cat(result_row, dim=3))
dec = torch.cat(result_rows, dim=2)
return dec
def forward(
self,
sample: torch.Tensor,
sample_posterior: bool = False,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
r"""
Args:
sample (`torch.Tensor`): Input sample.
sample_posterior (`bool`, *optional*, defaults to `False`):
Whether to sample from the posterior.
"""
x = sample
posterior = self.encode(x).latent_dist
if sample_posterior:
z = posterior.sample(generator=generator)
else:
z = posterior.mode()
dec = self.decode(z).sample
return dec
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections
def fuse_qkv_projections(self):
"""
Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value)
are fused. For cross-attention modules, key and value projection matrices are fused.
<Tip warning={true}>
This API is 🧪 experimental.
</Tip>
"""
self.original_attn_processors = None
for _, attn_processor in self.attn_processors.items():
if "Added" in str(attn_processor.__class__.__name__):
raise ValueError(
"`fuse_qkv_projections()` is not supported for models having added KV projections."
)
self.original_attn_processors = self.attn_processors
for module in self.modules():
if isinstance(module, Attention):
module.fuse_projections(fuse=True)
self.set_attn_processor(FusedAttnProcessor2_0())
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections
def unfuse_qkv_projections(self):
"""Disables the fused QKV projection if enabled.
<Tip warning={true}>
This API is 🧪 experimental.
</Tip>
"""
if self.original_attn_processors is not None:
self.set_attn_processor(self.original_attn_processors)
+5 -5
View File
@@ -10,7 +10,7 @@
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
from typing import Any
from typing import Any, List, Optional, Tuple
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: torch.Generator | None = None,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
"""
Args:
+11 -13
View File
@@ -16,6 +16,7 @@
import contextvars
from contextlib import contextmanager
from typing import Optional, Tuple, Union
import torch
import torch.nn as nn
@@ -67,9 +68,9 @@ class WanCausalConv3d(nn.Conv3d):
self,
in_channels: int,
out_channels: int,
kernel_size: int | tuple[int, int, int],
stride: int | tuple[int, int, int] = 1,
padding: int | tuple[int, int, int] = 0,
kernel_size: Union[int, Tuple[int, int, int]],
stride: Union[int, Tuple[int, int, int]] = 1,
padding: Union[int, Tuple[int, int, int]] = 0,
) -> None:
super().__init__(
in_channels=in_channels,
@@ -78,9 +79,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)
@@ -433,8 +434,7 @@ 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:],
strict=False):
for attn, resnet in zip(self.attentions, self.resnets[1:]):
if attn is not None:
x = attn(x)
@@ -488,8 +488,7 @@ class WanEncoder3d(nn.Module):
# downsample blocks
self.down_blocks = nn.ModuleList([])
for i, (in_dim,
out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=False)):
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
# residual (+attention) blocks
for _ in range(num_res_blocks):
self.down_blocks.append(
@@ -590,7 +589,7 @@ class WanUpBlock(nn.Module):
out_dim: int,
num_res_blocks: int,
dropout: float = 0.0,
upsample_mode: str | None = None,
upsample_mode: Optional[str] = None,
non_linearity: str = "silu",
):
super().__init__()
@@ -688,8 +687,7 @@ class WanDecoder3d(nn.Module):
# upsample blocks
self.up_blocks = nn.ModuleList([])
for i, (in_dim,
out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=False)):
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
# residual (+attention) blocks
if i > 0:
in_dim = in_dim // 2
@@ -946,7 +944,7 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
self,
sample: torch.Tensor,
sample_posterior: bool = False,
generator: torch.Generator | None = None,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
"""
Args:
+15 -11
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import os
from collections.abc import Callable
from typing import Callable, List, Optional, Tuple, Union
import numpy as np
import PIL.Image
@@ -29,7 +29,8 @@ else:
}
def pil_to_numpy(images: list[PIL.Image.Image] | PIL.Image.Image) -> np.ndarray:
def pil_to_numpy(
images: Union[List[PIL.Image.Image], PIL.Image.Image]) -> np.ndarray:
r"""
Convert a PIL image or a list of PIL images to NumPy arrays.
@@ -68,7 +69,9 @@ def numpy_to_pt(images: np.ndarray) -> torch.Tensor:
return images
def normalize(images: np.ndarray | torch.Tensor) -> np.ndarray | torch.Tensor:
def normalize(
images: Union[np.ndarray,
torch.Tensor]) -> Union[np.ndarray, torch.Tensor]:
r"""
Normalize an image array to [-1,1].
@@ -84,8 +87,9 @@ def normalize(images: np.ndarray | torch.Tensor) -> np.ndarray | torch.Tensor:
def load_image(
image: str | PIL.Image.Image,
convert_method: Callable[[PIL.Image.Image], PIL.Image.Image] | None = None
image: Union[str, PIL.Image.Image],
convert_method: Optional[Callable[[PIL.Image.Image],
PIL.Image.Image]] = None
) -> PIL.Image.Image:
"""
Loads `image` to a PIL Image.
@@ -128,11 +132,11 @@ def load_image(
def get_default_height_width(
image: PIL.Image.Image | np.ndarray | torch.Tensor,
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
vae_scale_factor: int,
height: int | None = None,
width: int | None = None,
) -> tuple[int, int]:
height: Optional[int] = None,
width: Optional[int] = None,
) -> Tuple[int, int]:
r"""
Returns the height and width of the image, downscaled to the next integer multiple of `vae_scale_factor`.
@@ -175,12 +179,12 @@ def get_default_height_width(
def resize(
image: PIL.Image.Image | np.ndarray | torch.Tensor,
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
height: int,
width: int,
resize_mode: str = "default", # "default", "fill", "crop"
resample: str = "lanczos",
) -> PIL.Image.Image | np.ndarray | torch.Tensor:
) -> Union[PIL.Image.Image, np.ndarray, torch.Tensor]:
"""
Resize image.
@@ -8,7 +8,7 @@ This module defines the base class for pipelines that are composed of multiple s
import os
from abc import ABC, abstractmethod
from copy import deepcopy
from typing import Any, cast
from typing import Any, Dict, List, Optional, cast
import torch
@@ -33,20 +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: dict[str, Any] | None = None):
config: Optional[Dict[str, Any]] = 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.model_path = model_path
self._stages: list[PipelineStage] = []
self._stage_name_mapping: dict[str, PipelineStage] = {}
self._stages: List[PipelineStage] = []
self._stage_name_mapping: Dict[str, PipelineStage] = {}
if self._required_config_modules is None:
raise NotImplementedError(
@@ -74,16 +74,16 @@ class ComposedPipelineBase(ABC):
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
@@ -101,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.
"""
@@ -120,7 +120,7 @@ class ComposedPipelineBase(ABC):
"""
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.
"""
@@ -128,8 +128,12 @@ class ComposedPipelineBase(ABC):
modules_config = deepcopy(self.config)
# remove keys that are not pipeline modules
modules_config.pop("_class_name")
modules_config.pop("_diffusers_version")
keys_to_remove = [
key for key in list(modules_config.keys())
if str(key).startswith("_")
]
for key in keys_to_remove:
modules_config.pop(key)
# some sanity checks
assert len(
@@ -0,0 +1,92 @@
import numpy as np
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
logger = init_logger(__name__)
class TimestepsPreparationPreStage(PipelineStage):
def __init__(self, scheduler) -> None:
super().__init__()
self.scheduler = scheduler
def calculate_shift(self,
image_seq_len,
base_seq_len: int = 256,
max_seq_len: int = 4096,
base_shift: float = 0.5,
max_shift: float = 1.15):
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - m * base_seq_len
mu = image_seq_len * m + b
return mu
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
assert batch.height is not None
assert batch.width is not None
batch.sigmas = np.linspace(
1.0, 1 / batch.num_inference_steps,
batch.num_inference_steps) if batch.sigmas is None else batch.sigmas
spatial_compression_ratio = fastvideo_args.vae_config.arch_config.spatial_compression_ratio
batch.extra_set_timesteps_kwargs["mu"] = (
batch.extra_set_timesteps_kwargs.get("mu", None)
or self.calculate_shift(
(batch.height // spatial_compression_ratio // 2) *
(batch.width // spatial_compression_ratio // 2),
self.scheduler.config.get("base_image_seq_len", 256),
self.scheduler.config.get("max_image_seq_len", 4096),
self.scheduler.config.get("base_shift", 0.5),
self.scheduler.config.get("max_shift", 1.15),
))
return batch
class DenoisingPreprocessingStage(PipelineStage):
def __init__(self) -> None:
super().__init__()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
# [B, in_channels // 4, 1, H, W] -> [B, H // 2, W // 2, in_channels]
assert batch.latents is not None
b, c, _, h, w = batch.latents.shape
latents = batch.latents.view(b, c, h // 2, 2, w // 2, 2)
latents = latents.permute(0, 2, 4, 1, 3, 5)
latents = latents.reshape(b, (h // 2) * (w // 2), c * 4)
batch.latents = latents
return batch
class DenoisingPostprocessingStage(PipelineStage):
def __init__(self) -> None:
super().__init__()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
assert batch.latents is not None
assert batch.height_latents is not None
assert batch.width_latents is not None
latents = batch.latents
# Skip decoding if output type is latent
if fastvideo_args.output_type == "latent":
latents = latents
else:
# [B, (H // 2) * (W // 2), in_channels] -> [B, in_channels // 4, 1, H, W]
b, _, c = latents.shape
h, w = batch.height_latents, batch.width_latents
latents = latents.view(b, h // 2, w // 2, c // 4, 2, 2)
latents = latents.permute(0, 3, 1, 4, 2, 5)
latents = latents.reshape(b, c // 4, h, w)
# latents = latents.squeeze(2)
batch.latents = latents
return batch
@@ -0,0 +1,80 @@
# SPDX-License-Identifier: Apache-2.0
"""
Flux image diffusion pipeline implementation.
This module contains an implementation of the Flux image diffusion pipeline
using the modular pipeline architecture.
"""
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.flux.custom_stages import (
DenoisingPostprocessingStage, DenoisingPreprocessingStage,
TimestepsPreparationPreStage)
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
logger = init_logger(__name__)
class FluxPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2"),
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2"),
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timesteps_preparation_pre_stage",
stage=TimestepsPreparationPreStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="modulation",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="denoising_preprocessing_stage",
stage=DenoisingPreprocessingStage())
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="denoising_postprocessing_stage",
stage=DenoisingPostprocessingStage())
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = FluxPipeline

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