Compare commits

...
Author SHA1 Message Date
SolitaryThinker d0c53871a4 fix tensor type hint 2025-05-23 14:42:23 -07:00
SolitaryThinker 0d5306f61f update min python to 3.10 2025-05-23 14:42:23 -07:00
91 changed files with 814 additions and 853 deletions
+1 -3
View File
@@ -18,7 +18,6 @@ import os
import re import re
import sys import sys
from pathlib import Path from pathlib import Path
from typing import Optional
import requests import requests
@@ -168,8 +167,7 @@ _cached_base: str = ""
_cached_branch: str = "" _cached_branch: str = ""
def get_repo_base_and_branch( def get_repo_base_and_branch(pr_number: str) -> tuple[str | None, str | None]:
pr_number: str) -> tuple[Optional[str], Optional[str]]:
global _cached_base, _cached_branch global _cached_base, _cached_branch
if _cached_base and _cached_branch: if _cached_base and _cached_branch:
return _cached_base, _cached_branch return _cached_base, _cached_branch
+1 -2
View File
@@ -5,7 +5,6 @@ import itertools
import re import re
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Optional
ROOT_DIR = Path(__file__).parent.parent.parent.resolve() ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
ROOT_DIR_RELATIVE = '../../../..' ROOT_DIR_RELATIVE = '../../../..'
@@ -89,7 +88,7 @@ class Example:
generate() -> str: Generates the documentation content. generate() -> str: Generates the documentation content.
""" # noqa: E501 """ # noqa: E501
path: Path path: Path
category: Optional[str] = None category: str | None = None
main_file: Path = field(init=False) main_file: Path = field(init=False)
other_files: list[Path] = field(init=False) other_files: list[Path] = field(init=False)
title: str = field(init=False) title: str = field(init=False)
+6 -8
View File
@@ -3,8 +3,7 @@
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from dataclasses import dataclass, fields from dataclasses import dataclass, fields
from typing import (TYPE_CHECKING, Any, Dict, Generic, Optional, Protocol, Set, from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar
Type, TypeVar)
if TYPE_CHECKING: if TYPE_CHECKING:
from fastvideo.v1.fastvideo_args import FastVideoArgs from fastvideo.v1.fastvideo_args import FastVideoArgs
@@ -27,12 +26,12 @@ class AttentionBackend(ABC):
@staticmethod @staticmethod
@abstractmethod @abstractmethod
def get_impl_cls() -> Type["AttentionImpl"]: def get_impl_cls() -> type["AttentionImpl"]:
raise NotImplementedError raise NotImplementedError
@staticmethod @staticmethod
@abstractmethod @abstractmethod
def get_metadata_cls() -> Type["AttentionMetadata"]: def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError raise NotImplementedError
# @staticmethod # @staticmethod
@@ -46,7 +45,7 @@ class AttentionBackend(ABC):
@staticmethod @staticmethod
@abstractmethod @abstractmethod
def get_builder_cls() -> Type["AttentionMetadataBuilder"]: def get_builder_cls() -> type["AttentionMetadataBuilder"]:
raise NotImplementedError raise NotImplementedError
@@ -57,8 +56,7 @@ class AttentionMetadata:
current_timestep: int current_timestep: int
def asdict_zerocopy(self, def asdict_zerocopy(self,
skip_fields: Optional[Set[str]] = None skip_fields: set[str] | None = None) -> dict[str, Any]:
) -> Dict[str, Any]:
"""Similar to dataclasses.asdict, but avoids deepcopying.""" """Similar to dataclasses.asdict, but avoids deepcopying."""
if skip_fields is None: if skip_fields is None:
skip_fields = set() skip_fields = set()
@@ -124,7 +122,7 @@ class AttentionImpl(ABC, Generic[T]):
head_size: int, head_size: int,
softmax_scale: float, softmax_scale: float,
causal: bool = False, causal: bool = False,
num_kv_heads: Optional[int] = None, num_kv_heads: int | None = None,
prefix: str = "", prefix: str = "",
**extra_impl_args, **extra_impl_args,
) -> None: ) -> None:
@@ -1,7 +1,5 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from typing import List, Optional, Type
import torch import torch
from flash_attn import flash_attn_func as flash_attn_2_func from flash_attn import flash_attn_func as flash_attn_2_func
@@ -28,7 +26,7 @@ class FlashAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True accept_output_buffer: bool = True
@staticmethod @staticmethod
def get_supported_head_sizes() -> List[int]: def get_supported_head_sizes() -> list[int]:
return [32, 64, 96, 128, 160, 192, 224, 256] return [32, 64, 96, 128, 160, 192, 224, 256]
@staticmethod @staticmethod
@@ -36,15 +34,15 @@ class FlashAttentionBackend(AttentionBackend):
return "FLASH_ATTN" return "FLASH_ATTN"
@staticmethod @staticmethod
def get_impl_cls() -> Type["FlashAttentionImpl"]: def get_impl_cls() -> type["FlashAttentionImpl"]:
return FlashAttentionImpl return FlashAttentionImpl
@staticmethod @staticmethod
def get_metadata_cls() -> Type["AttentionMetadata"]: def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError raise NotImplementedError
@staticmethod @staticmethod
def get_builder_cls() -> Type["AttentionMetadataBuilder"]: def get_builder_cls() -> type["AttentionMetadataBuilder"]:
raise NotImplementedError raise NotImplementedError
@@ -56,7 +54,7 @@ class FlashAttentionImpl(AttentionImpl):
head_size: int, head_size: int,
causal: bool, causal: bool,
softmax_scale: float, softmax_scale: float,
num_kv_heads: Optional[int] = None, num_kv_heads: int | None = None,
prefix: str = "", prefix: str = "",
**extra_impl_args, **extra_impl_args,
) -> None: ) -> None:
+3 -5
View File
@@ -1,5 +1,3 @@
from typing import List, Optional, Type
import torch import torch
from sageattention import sageattn from sageattention import sageattn
@@ -17,7 +15,7 @@ class SageAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True accept_output_buffer: bool = True
@staticmethod @staticmethod
def get_supported_head_sizes() -> List[int]: def get_supported_head_sizes() -> list[int]:
return [32, 64, 96, 128, 160, 192, 224, 256] return [32, 64, 96, 128, 160, 192, 224, 256]
@staticmethod @staticmethod
@@ -25,7 +23,7 @@ class SageAttentionBackend(AttentionBackend):
return "SAGE_ATTN" return "SAGE_ATTN"
@staticmethod @staticmethod
def get_impl_cls() -> Type["SageAttentionImpl"]: def get_impl_cls() -> type["SageAttentionImpl"]:
return SageAttentionImpl return SageAttentionImpl
# @staticmethod # @staticmethod
@@ -41,7 +39,7 @@ class SageAttentionImpl(AttentionImpl):
head_size: int, head_size: int,
causal: bool, causal: bool,
softmax_scale: float, softmax_scale: float,
num_kv_heads: Optional[int] = None, num_kv_heads: int | None = None,
prefix: str = "", prefix: str = "",
**extra_impl_args, **extra_impl_args,
) -> None: ) -> None:
+3 -5
View File
@@ -1,5 +1,3 @@
from typing import List, Optional, Type
import torch import torch
from fastvideo.v1.attention.backends.abstract import ( from fastvideo.v1.attention.backends.abstract import (
@@ -16,7 +14,7 @@ class SDPABackend(AttentionBackend):
accept_output_buffer: bool = True accept_output_buffer: bool = True
@staticmethod @staticmethod
def get_supported_head_sizes() -> List[int]: def get_supported_head_sizes() -> list[int]:
return [32, 64, 96, 128, 160, 192, 224, 256] return [32, 64, 96, 128, 160, 192, 224, 256]
@staticmethod @staticmethod
@@ -24,7 +22,7 @@ class SDPABackend(AttentionBackend):
return "SDPA" return "SDPA"
@staticmethod @staticmethod
def get_impl_cls() -> Type["SDPAImpl"]: def get_impl_cls() -> type["SDPAImpl"]:
return SDPAImpl return SDPAImpl
# @staticmethod # @staticmethod
@@ -40,7 +38,7 @@ class SDPAImpl(AttentionImpl):
head_size: int, head_size: int,
causal: bool, causal: bool,
softmax_scale: float, softmax_scale: float,
num_kv_heads: Optional[int] = None, num_kv_heads: int | None = None,
prefix: str = "", prefix: str = "",
**extra_impl_args, **extra_impl_args,
) -> None: ) -> None:
@@ -1,6 +1,5 @@
import json import json
from dataclasses import dataclass from dataclasses import dataclass
from typing import List, Optional, Type
import torch import torch
from einops import rearrange from einops import rearrange
@@ -20,7 +19,7 @@ logger = init_logger(__name__)
# TODO(will-refactor): move this to a utils file # TODO(will-refactor): move this to a utils file
def dict_to_3d_list(mask_strategy) -> List[List[List[Optional[torch.Tensor]]]]: def dict_to_3d_list(mask_strategy) -> list[list[list[torch.Tensor | None]]]:
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy] indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
max_timesteps_idx = max( max_timesteps_idx = max(
@@ -58,7 +57,7 @@ class SlidingTileAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True accept_output_buffer: bool = True
@staticmethod @staticmethod
def get_supported_head_sizes() -> List[int]: def get_supported_head_sizes() -> list[int]:
# TODO(will-refactor): check this # TODO(will-refactor): check this
return [32, 64, 96, 128, 160, 192, 224, 256] return [32, 64, 96, 128, 160, 192, 224, 256]
@@ -67,15 +66,15 @@ class SlidingTileAttentionBackend(AttentionBackend):
return "SLIDING_TILE_ATTN" return "SLIDING_TILE_ATTN"
@staticmethod @staticmethod
def get_impl_cls() -> Type["SlidingTileAttentionImpl"]: def get_impl_cls() -> type["SlidingTileAttentionImpl"]:
return SlidingTileAttentionImpl return SlidingTileAttentionImpl
@staticmethod @staticmethod
def get_metadata_cls() -> Type["SlidingTileAttentionMetadata"]: def get_metadata_cls() -> type["SlidingTileAttentionMetadata"]:
return SlidingTileAttentionMetadata return SlidingTileAttentionMetadata
@staticmethod @staticmethod
def get_builder_cls() -> Type["SlidingTileAttentionMetadataBuilder"]: def get_builder_cls() -> type["SlidingTileAttentionMetadataBuilder"]:
return SlidingTileAttentionMetadataBuilder return SlidingTileAttentionMetadataBuilder
@@ -110,7 +109,7 @@ class SlidingTileAttentionImpl(AttentionImpl):
head_size: int, head_size: int,
causal: bool, causal: bool,
softmax_scale: float, softmax_scale: float,
num_kv_heads: Optional[int] = None, num_kv_heads: int | None = None,
prefix: str = "", prefix: str = "",
**extra_impl_args, **extra_impl_args,
) -> None: ) -> None:
+12 -14
View File
@@ -1,7 +1,5 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from typing import Optional, Tuple
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -22,11 +20,11 @@ class DistributedAttention(nn.Module):
def __init__(self, def __init__(self,
num_heads: int, num_heads: int,
head_size: int, head_size: int,
num_kv_heads: Optional[int] = None, num_kv_heads: int | None = None,
softmax_scale: Optional[float] = None, softmax_scale: float | None = None,
causal: bool = False, causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend, supported_attention_backends: tuple[_Backend, ...]
...]] = None, | None = None,
prefix: str = "", prefix: str = "",
**extra_impl_args) -> None: **extra_impl_args) -> None:
super().__init__() super().__init__()
@@ -62,10 +60,10 @@ class DistributedAttention(nn.Module):
q: torch.Tensor, q: torch.Tensor,
k: torch.Tensor, k: torch.Tensor,
v: torch.Tensor, v: torch.Tensor,
replicated_q: Optional[torch.Tensor] = None, replicated_q: torch.Tensor | None = None,
replicated_k: Optional[torch.Tensor] = None, replicated_k: torch.Tensor | None = None,
replicated_v: Optional[torch.Tensor] = None, replicated_v: torch.Tensor | None = None,
) -> tuple[torch.Tensor, Optional[torch.Tensor]]: ) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Forward pass for distributed attention. """Forward pass for distributed attention.
Args: Args:
@@ -141,11 +139,11 @@ class LocalAttention(nn.Module):
def __init__(self, def __init__(self,
num_heads: int, num_heads: int,
head_size: int, head_size: int,
num_kv_heads: Optional[int] = None, num_kv_heads: int | None = None,
softmax_scale: Optional[float] = None, softmax_scale: float | None = None,
causal: bool = False, causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend, supported_attention_backends: tuple[_Backend, ...]
...]] = None, | None = None,
**extra_impl_args) -> None: **extra_impl_args) -> None:
super().__init__() super().__init__()
if softmax_scale is None: if softmax_scale is None:
+14 -13
View File
@@ -2,9 +2,10 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/selector.py # Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/selector.py
import os import os
from collections.abc import Generator
from contextlib import contextmanager from contextlib import contextmanager
from functools import cache from functools import cache
from typing import Generator, Optional, Tuple, Type, cast from typing import cast
import torch import torch
@@ -17,7 +18,7 @@ from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
logger = init_logger(__name__) logger = init_logger(__name__)
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]: def backend_name_to_enum(backend_name: str) -> _Backend | None:
""" """
Convert a string backend name to a _Backend enum value. Convert a string backend name to a _Backend enum value.
@@ -31,7 +32,7 @@ def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
None None
def get_env_variable_attn_backend() -> Optional[_Backend]: def get_env_variable_attn_backend() -> _Backend | None:
''' '''
Get the backend override specified by the FastVideo attention Get the backend override specified by the FastVideo attention
backend environment variable, if one is specified. backend environment variable, if one is specified.
@@ -53,10 +54,10 @@ def get_env_variable_attn_backend() -> Optional[_Backend]:
# #
# THIS SELECTION TAKES PRECEDENCE OVER THE # THIS SELECTION TAKES PRECEDENCE OVER THE
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE # FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
forced_attn_backend: Optional[_Backend] = None forced_attn_backend: _Backend | None = None
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None: def global_force_attn_backend(attn_backend: _Backend | None) -> None:
''' '''
Force all attention operations to use a specified backend. Force all attention operations to use a specified backend.
@@ -71,7 +72,7 @@ def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
forced_attn_backend = attn_backend forced_attn_backend = attn_backend
def get_global_forced_attn_backend() -> Optional[_Backend]: def get_global_forced_attn_backend() -> _Backend | None:
''' '''
Get the currently-forced choice of attention backend, Get the currently-forced choice of attention backend,
or None if auto-selection is currently enabled. or None if auto-selection is currently enabled.
@@ -82,8 +83,8 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_attn_backend( def get_attn_backend(
head_size: int, head_size: int,
dtype: torch.dtype, dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None, supported_attention_backends: tuple[_Backend, ...] | None = None,
) -> Type[AttentionBackend]: ) -> type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype, return _cached_get_attn_backend(head_size, dtype,
supported_attention_backends) supported_attention_backends)
@@ -92,8 +93,8 @@ def get_attn_backend(
def _cached_get_attn_backend( def _cached_get_attn_backend(
head_size: int, head_size: int,
dtype: torch.dtype, dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None, supported_attention_backends: tuple[_Backend, ...] | None = None,
) -> Type[AttentionBackend]: ) -> type[AttentionBackend]:
# Check whether a particular choice of backend was # Check whether a particular choice of backend was
# previously forced. # previously forced.
# #
@@ -102,13 +103,13 @@ def _cached_get_attn_backend(
if not supported_attention_backends: if not supported_attention_backends:
raise ValueError("supported_attention_backends is empty") raise ValueError("supported_attention_backends is empty")
selected_backend = None selected_backend = None
backend_by_global_setting: Optional[_Backend] = ( backend_by_global_setting: _Backend | None = (
get_global_forced_attn_backend()) get_global_forced_attn_backend())
if backend_by_global_setting is not None: if backend_by_global_setting is not None:
selected_backend = backend_by_global_setting selected_backend = backend_by_global_setting
else: else:
# Check the environment variable and override if specified # Check the environment variable and override if specified
backend_by_env_var: Optional[str] = envs.FASTVIDEO_ATTENTION_BACKEND backend_by_env_var: str | None = envs.FASTVIDEO_ATTENTION_BACKEND
if backend_by_env_var is not None: if backend_by_env_var is not None:
selected_backend = backend_name_to_enum(backend_by_env_var) selected_backend = backend_name_to_enum(backend_by_env_var)
@@ -120,7 +121,7 @@ def _cached_get_attn_backend(
if not attention_cls: if not attention_cls:
raise ValueError( raise ValueError(
f"Invalid attention backend for {current_platform.device_name}") 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 @contextmanager
+3 -3
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field, fields from dataclasses import dataclass, field, fields
from typing import Any, Dict from typing import Any
from fastvideo.v1.logger import init_logger from fastvideo.v1.logger import init_logger
@@ -41,7 +41,7 @@ class ModelConfig:
self.__dict__.update(state) self.__dict__.update(state)
# This should be used only when loading from transformers/diffusers # This should be used only when loading from transformers/diffusers
def update_model_arch(self, source_model_dict: Dict[str, Any]) -> None: def update_model_arch(self, source_model_dict: dict[str, Any]) -> None:
arch_config = self.arch_config arch_config = self.arch_config
valid_fields = {f.name for f in fields(arch_config)} valid_fields = {f.name for f in fields(arch_config)}
@@ -55,7 +55,7 @@ class ModelConfig:
if hasattr(arch_config, "__post_init__"): if hasattr(arch_config, "__post_init__"):
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." 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)} valid_fields = {f.name for f in fields(self)}
+3 -3
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any, Optional, Tuple from typing import Any
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.layers.quantization import QuantizationConfig from fastvideo.v1.layers.quantization import QuantizationConfig
@@ -11,7 +11,7 @@ class DiTArchConfig(ArchConfig):
_fsdp_shard_conditions: list = field(default_factory=list) _fsdp_shard_conditions: list = field(default_factory=list)
_compile_conditions: list = field(default_factory=list) _compile_conditions: list = field(default_factory=list)
_param_names_mapping: dict = field(default_factory=dict) _param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: Tuple[_Backend, _supported_attention_backends: tuple[_Backend,
...] = (_Backend.SLIDING_TILE_ATTN, ...] = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_ATTN, _Backend.SAGE_ATTN,
_Backend.FLASH_ATTN, _Backend.FLASH_ATTN,
@@ -32,7 +32,7 @@ class DiTConfig(ModelConfig):
# FastVideoDiT-specific parameters # FastVideoDiT-specific parameters
prefix: str = "" prefix: str = ""
quant_config: Optional[QuantizationConfig] = None quant_config: QuantizationConfig | None = None
@staticmethod @staticmethod
def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any: def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any:
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Optional, Tuple
import torch import torch
@@ -156,9 +155,9 @@ class HunyuanVideoArchConfig(DiTArchConfig):
num_layers: int = 20 num_layers: int = 20
num_single_layers: int = 40 num_single_layers: int = 40
num_refiner_layers: int = 2 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 guidance_embeds: bool = False
dtype: Optional[torch.dtype] = None dtype: torch.dtype | None = None
text_embed_dim: int = 4096 text_embed_dim: int = 4096
pooled_projection_dim: int = 768 pooled_projection_dim: int = 768
rope_theta: int = 256 rope_theta: int = 256
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import List, Optional, Tuple, Union
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -40,17 +39,17 @@ class StepVideoArchConfig(DiTArchConfig):
num_attention_heads: int = 48 num_attention_heads: int = 48
attention_head_dim: int = 128 attention_head_dim: int = 128
in_channels: int = 64 in_channels: int = 64
out_channels: Optional[int] = 64 out_channels: int | None = 64
num_layers: int = 48 num_layers: int = 48
dropout: float = 0.0 dropout: float = 0.0
patch_size: int = 1 patch_size: int = 1
norm_type: str = "ada_norm_single" norm_type: str = "ada_norm_single"
norm_elementwise_affine: bool = False norm_elementwise_affine: bool = False
norm_eps: float = 1e-6 norm_eps: float = 1e-6
caption_channels: Optional[Union[int, List[int], Tuple[int, ...]]] = field( caption_channels: int | list[int] | tuple[int, ...] | None = field(
default_factory=lambda: [6144, 1024]) default_factory=lambda: [6144, 1024])
attention_type: Optional[str] = "torch" attention_type: str | None = "torch"
use_additional_conditions: Optional[bool] = False use_additional_conditions: bool | None = False
def __post_init__(self): def __post_init__(self):
self.hidden_size = self.num_attention_heads * self.attention_head_dim self.hidden_size = self.num_attention_heads * self.attention_head_dim
+3 -4
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Optional, Tuple
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -52,7 +51,7 @@ class WanVideoArchConfig(DiTArchConfig):
r"blocks.\1.self_attn_residual_norm.norm.\2", 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 text_len = 512
num_attention_heads: int = 40 num_attention_heads: int = 40
attention_head_dim: int = 128 attention_head_dim: int = 128
@@ -65,8 +64,8 @@ class WanVideoArchConfig(DiTArchConfig):
cross_attn_norm: bool = True cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads" qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6 eps: float = 1e-6
image_dim: Optional[int] = None image_dim: int | None = None
added_kv_proj_dim: Optional[int] = None added_kv_proj_dim: int | None = None
rope_max_seq_len: int = 1024 rope_max_seq_len: int = 1024
def __post_init__(self): def __post_init__(self):
+11 -11
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Tuple from typing import Any
import torch import torch
@@ -10,8 +10,8 @@ from fastvideo.v1.platforms import _Backend
@dataclass @dataclass
class EncoderArchConfig(ArchConfig): class EncoderArchConfig(ArchConfig):
architectures: List[str] = field(default_factory=lambda: []) architectures: list[str] = field(default_factory=lambda: [])
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN, _supported_attention_backends: tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) _Backend.TORCH_SDPA)
output_hidden_states: bool = False output_hidden_states: bool = False
use_return_dict: bool = True use_return_dict: bool = True
@@ -32,7 +32,7 @@ class TextEncoderArchConfig(EncoderArchConfig):
scalable_attention: bool = True scalable_attention: bool = True
tie_word_embeddings: bool = False 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: def __post_init__(self) -> None:
self.tokenizer_kwargs = { self.tokenizer_kwargs = {
@@ -49,11 +49,11 @@ class ImageEncoderArchConfig(EncoderArchConfig):
@dataclass @dataclass
class BaseEncoderOutput: class BaseEncoderOutput:
last_hidden_state: Optional[torch.FloatTensor] = None last_hidden_state: torch.FloatTensor | None = None
pooler_output: Optional[torch.FloatTensor] = None pooler_output: torch.FloatTensor | None = None
hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None hidden_states: tuple[torch.FloatTensor, ...] | None = None
attentions: Optional[Tuple[torch.FloatTensor, ...]] = None attentions: tuple[torch.FloatTensor, ...] | None = None
attention_mask: Optional[torch.Tensor] = None attention_mask: torch.Tensor | None = None
@dataclass @dataclass
@@ -61,8 +61,8 @@ class EncoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=EncoderArchConfig) arch_config: ArchConfig = field(default_factory=EncoderArchConfig)
prefix: str = "" prefix: str = ""
quant_config: Optional[QuantizationConfig] = None quant_config: QuantizationConfig | None = None
lora_config: Optional[Any] = None lora_config: Any | None = None
@dataclass @dataclass
+4 -5
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig, from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
ImageEncoderConfig, ImageEncoderConfig,
@@ -51,8 +50,8 @@ class CLIPTextConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field( arch_config: TextEncoderArchConfig = field(
default_factory=CLIPTextArchConfig) default_factory=CLIPTextArchConfig)
num_hidden_layers_override: Optional[int] = None num_hidden_layers_override: int | None = None
require_post_norm: Optional[bool] = None require_post_norm: bool | None = None
prefix: str = "clip" prefix: str = "clip"
@@ -61,6 +60,6 @@ class CLIPVisionConfig(ImageEncoderConfig):
arch_config: ImageEncoderArchConfig = field( arch_config: ImageEncoderArchConfig = field(
default_factory=CLIPVisionArchConfig) default_factory=CLIPVisionArchConfig)
num_hidden_layers_override: Optional[int] = None num_hidden_layers_override: int | None = None
require_post_norm: Optional[bool] = None require_post_norm: bool | None = None
prefix: str = "clip" prefix: str = "clip"
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig, from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig) TextEncoderConfig)
@@ -12,7 +11,7 @@ class LlamaArchConfig(TextEncoderArchConfig):
intermediate_size: int = 11008 intermediate_size: int = 11008
num_hidden_layers: int = 32 num_hidden_layers: int = 32
num_attention_heads: int = 32 num_attention_heads: int = 32
num_key_value_heads: Optional[int] = None num_key_value_heads: int | None = None
hidden_act: str = "silu" hidden_act: str = "silu"
max_position_embeddings: int = 2048 max_position_embeddings: int = 2048
initializer_range: float = 0.02 initializer_range: float = 0.02
@@ -24,11 +23,11 @@ class LlamaArchConfig(TextEncoderArchConfig):
pretraining_tp: int = 1 pretraining_tp: int = 1
tie_word_embeddings: bool = False tie_word_embeddings: bool = False
rope_theta: float = 10000.0 rope_theta: float = 10000.0
rope_scaling: Optional[float] = None rope_scaling: float | None = None
attention_bias: bool = False attention_bias: bool = False
attention_dropout: float = 0.0 attention_dropout: float = 0.0
mlp_bias: bool = False mlp_bias: bool = False
head_dim: Optional[int] = None head_dim: int | None = None
hidden_state_skip_layer: int = 2 hidden_state_skip_layer: int = 2
text_len: int = 256 text_len: int = 256
+1 -2
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig, from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig) TextEncoderConfig)
@@ -12,7 +11,7 @@ class T5ArchConfig(TextEncoderArchConfig):
d_kv: int = 64 d_kv: int = 64
d_ff: int = 2048 d_ff: int = 2048
num_layers: int = 6 num_layers: int = 6
num_decoder_layers: Optional[int] = None num_decoder_layers: int | None = None
num_heads: int = 8 num_heads: int = 8
relative_attention_num_buckets: int = 32 relative_attention_num_buckets: int = 32
relative_attention_max_distance: int = 128 relative_attention_max_distance: int = 128
+2 -2
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any, Union from typing import Any
import torch import torch
@@ -9,7 +9,7 @@ from fastvideo.v1.utils import StoreBoolean
@dataclass @dataclass
class VAEArchConfig(ArchConfig): class VAEArchConfig(ArchConfig):
scaling_factor: Union[float, torch.tensor] = 0 scaling_factor: float | torch.Tensor = 0
temporal_compression_ratio: int = 4 temporal_compression_ratio: int = 4
spatial_compression_ratio: int = 8 spatial_compression_ratio: int = 8
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Tuple
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
@@ -9,19 +8,19 @@ class HunyuanVAEArchConfig(VAEArchConfig):
in_channels: int = 3 in_channels: int = 3
out_channels: int = 3 out_channels: int = 3
latent_channels: int = 16 latent_channels: int = 16
down_block_types: Tuple[str, ...] = ( down_block_types: tuple[str, ...] = (
"HunyuanVideoDownBlock3D", "HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D", "HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D", "HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D", "HunyuanVideoDownBlock3D",
) )
up_block_types: Tuple[str, ...] = ( up_block_types: tuple[str, ...] = (
"HunyuanVideoUpBlock3D", "HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D", "HunyuanVideoUpBlock3D",
"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 layers_per_block: int = 2
act_fn: str = "silu" act_fn: str = "silu"
norm_num_groups: int = 32 norm_num_groups: int = 32
+7 -8
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Tuple
import torch import torch
@@ -10,12 +9,12 @@ from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
class WanVAEArchConfig(VAEArchConfig): class WanVAEArchConfig(VAEArchConfig):
base_dim: int = 96 base_dim: int = 96
z_dim: int = 16 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 num_res_blocks: int = 2
attn_scales: Tuple[float, ...] = () attn_scales: tuple[float, ...] = ()
temperal_downsample: Tuple[bool, ...] = (False, True, True) temperal_downsample: tuple[bool, ...] = (False, True, True)
dropout: float = 0.0 dropout: float = 0.0
latents_mean: Tuple[float, ...] = ( latents_mean: tuple[float, ...] = (
-0.7571, -0.7571,
-0.7089, -0.7089,
-0.9113, -0.9113,
@@ -33,7 +32,7 @@ class WanVAEArchConfig(VAEArchConfig):
0.2503, 0.2503,
-0.2921, -0.2921,
) )
latents_std: Tuple[float, ...] = ( latents_std: tuple[float, ...] = (
2.8184, 2.8184,
1.4541, 1.4541,
2.3275, 2.3275,
@@ -55,9 +54,9 @@ class WanVAEArchConfig(VAEArchConfig):
spatial_compression_ratio = 8 spatial_compression_ratio = 8
def __post_init__(self): 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.latents_std).view(1, self.z_dim, 1, 1, 1)
self.shift_factor: torch.tensor = torch.tensor(self.latents_mean).view( self.shift_factor: torch.Tensor = torch.tensor(self.latents_mean).view(
1, self.z_dim, 1, 1, 1) 1, self.z_dim, 1, 1, 1)
+13 -11
View File
@@ -1,6 +1,7 @@
import json import json
from collections.abc import Callable
from dataclasses import asdict, dataclass, field, fields from dataclasses import asdict, dataclass, field, fields
from typing import Any, Callable, Dict, Optional, Tuple, cast from typing import Any, cast
import torch import torch
@@ -17,7 +18,7 @@ def preprocess_text(prompt: str) -> str:
return prompt return prompt
def postprocess_text(output: BaseEncoderOutput) -> torch.tensor: def postprocess_text(output: BaseEncoderOutput) -> torch.Tensor:
raise NotImplementedError raise NotImplementedError
@@ -26,7 +27,7 @@ class PipelineConfig:
"""Base configuration for all pipeline architectures.""" """Base configuration for all pipeline architectures."""
# Video generation parameters # Video generation parameters
embedded_cfg_scale: float = 6.0 embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None flow_shift: float | None = None
use_cpu_offload: bool = False use_cpu_offload: bool = False
disable_autocast: bool = False disable_autocast: bool = False
@@ -43,18 +44,18 @@ class PipelineConfig:
dit_config: DiTConfig = field(default_factory=DiTConfig) dit_config: DiTConfig = field(default_factory=DiTConfig)
# Text encoder configuration # Text encoder configuration
text_encoder_precisions: Tuple[str, ...] = field( text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp16", )) default_factory=lambda: ("fp16", ))
text_encoder_configs: Tuple[EncoderConfig, ...] = field( text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), )) 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, )) 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: ...] = field(default_factory=lambda:
(postprocess_text, )) (postprocess_text, ))
# STA (Spatial-Temporal Attention) parameters # STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None mask_strategy_file_path: str | None = None
# Compilation # Compilation
enable_torch_compile: bool = False enable_torch_compile: bool = False
@@ -107,7 +108,7 @@ class PipelineConfig:
input_pipeline_dict = json.load(f) input_pipeline_dict = json.load(f)
self.update_pipeline_config(input_pipeline_dict) 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: Any]) -> None:
for f in fields(self): for f in fields(self):
key = f.name key = f.name
@@ -123,8 +124,9 @@ class PipelineConfig:
assert len(current_value) == len( assert len(current_value) == len(
new_value new_value
), "Users shouldn't delete or add text encoder config objects in your json" ), "Users shouldn't delete or add text encoder config objects in your json"
for target_config, source_config in zip( for target_config, source_config in zip(current_value,
current_value, new_value): new_value,
strict=False):
target_config.update_model_config(source_config) target_config.update_model_config(source_config)
else: else:
setattr(self, key, new_value) setattr(self, key, new_value)
+12 -11
View File
@@ -1,5 +1,6 @@
from collections.abc import Callable
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Callable, Tuple, TypedDict from typing import TypedDict
import torch import torch
@@ -35,11 +36,11 @@ def llama_preprocess_text(prompt: str) -> str:
return prompt_template_video["template"].format(prompt) 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 hidden_state_skip_layer = 2
assert outputs.hidden_states is not None assert outputs.hidden_states is not None
hidden_states: Tuple[torch.Tensor, ...] = outputs.hidden_states hidden_states: tuple[torch.Tensor, ...] = outputs.hidden_states
last_hidden_state: torch.tensor = hidden_states[-(hidden_state_skip_layer + last_hidden_state: torch.Tensor = hidden_states[-(hidden_state_skip_layer +
1)] 1)]
crop_start = prompt_template_video.get("crop_start", -1) crop_start = prompt_template_video.get("crop_start", -1)
last_hidden_state = last_hidden_state[:, crop_start:] last_hidden_state = last_hidden_state[:, crop_start:]
@@ -50,8 +51,8 @@ def clip_preprocess_text(prompt: str) -> str:
return prompt return prompt
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor: def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
pooler_output: torch.tensor = outputs.pooler_output pooler_output: torch.Tensor = outputs.pooler_output
return pooler_output return pooler_output
@@ -72,19 +73,19 @@ class HunyuanConfig(PipelineConfig):
use_cpu_offload: bool = True use_cpu_offload: bool = True
# Text encoding stage # Text encoding stage
text_encoder_configs: Tuple[EncoderConfig, ...] = field( text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LlamaConfig(), CLIPTextConfig())) 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)) default_factory=lambda: (llama_preprocess_text, clip_preprocess_text))
postprocess_text_funcs: Tuple[ postprocess_text_funcs: tuple[
Callable[[BaseEncoderOutput], torch.tensor], Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda: ...] = field(default_factory=lambda:
(llama_postprocess_text, clip_postprocess_text)) (llama_postprocess_text, clip_postprocess_text))
# Precision for each component # Precision for each component
precision: str = "bf16" precision: str = "bf16"
vae_precision: str = "fp16" vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field( text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp16")) default_factory=lambda: ("fp16", "fp16"))
def __post_init__(self): def __post_init__(self):
+5 -5
View File
@@ -1,7 +1,7 @@
"""Registry for pipeline weight-specific configurations.""" """Registry for pipeline weight-specific configurations."""
import os import os
from typing import Callable, Dict, Optional, Type from collections.abc import Callable
from fastvideo.v1.configs.pipelines.base import PipelineConfig from fastvideo.v1.configs.pipelines.base import PipelineConfig
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig, from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
@@ -18,7 +18,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
logger = init_logger(__name__) logger = init_logger(__name__)
# Registry maps specific model weights to their config classes # 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, "FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig, "hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig, "Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
@@ -30,7 +30,7 @@ WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
} }
# For determining pipeline type from model ID # 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(), "hunyuan": lambda id: "hunyuan" in id.lower(),
"wanpipeline": lambda id: "wanpipeline" in id.lower(), "wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(), "wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
@@ -39,7 +39,7 @@ PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
} }
# Fallback configs when exact match isn't found but architecture is detected # 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": "hunyuan":
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
"wanpipeline": "wanpipeline":
@@ -51,7 +51,7 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
def get_pipeline_config_cls_for_name( def get_pipeline_config_cls_for_name(
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]: pipeline_name_or_path: str) -> type[PipelineConfig] | None:
"""Get the appropriate config class for specific pretrained weights.""" """Get the appropriate config class for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path): if os.path.exists(pipeline_name_or_path):
+11 -9
View File
@@ -1,5 +1,5 @@
from collections.abc import Callable
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Callable, Tuple
import torch import torch
@@ -11,13 +11,15 @@ from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig from fastvideo.v1.configs.pipelines.base import PipelineConfig
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor: def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
mask: torch.tensor = outputs.attention_mask mask: torch.Tensor = outputs.attention_mask
hidden_state: torch.tensor = outputs.last_hidden_state hidden_state: torch.Tensor = outputs.last_hidden_state
seq_lens = mask.gt(0).sum(dim=1).long() seq_lens = mask.gt(0).sum(dim=1).long()
assert torch.isnan(hidden_state).sum() == 0 assert torch.isnan(hidden_state).sum() == 0
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens)] prompt_embeds = [
prompt_embeds_tensor: torch.tensor = torch.stack([ u[:v] for u, v in zip(hidden_state, seq_lens, strict=False)
]
prompt_embeds_tensor: torch.Tensor = torch.stack([
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))]) torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
for u in prompt_embeds for u in prompt_embeds
], ],
@@ -44,16 +46,16 @@ class WanT2V480PConfig(PipelineConfig):
flow_shift: int = 3 flow_shift: int = 3
# Text encoding stage # Text encoding stage
text_encoder_configs: Tuple[EncoderConfig, ...] = field( text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (T5Config(), )) default_factory=lambda: (T5Config(), ))
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor], postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda: ...] = field(default_factory=lambda:
(t5_postprocess_text, )) (t5_postprocess_text, ))
# Precision for each component # Precision for each component
precision: str = "bf16" precision: str = "bf16"
vae_precision: str = "fp16" vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field( text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp32", )) default_factory=lambda: ("fp32", ))
# WanConfig-specific added parameters # WanConfig-specific added parameters
+6 -6
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Union from typing import Any
from fastvideo.v1.logger import init_logger from fastvideo.v1.logger import init_logger
@@ -15,12 +15,12 @@ class SamplingParam:
data_type: str = "video" data_type: str = "video"
# Image inputs # Image inputs
image_path: Optional[str] = None image_path: str | None = None
# Text inputs # Text inputs
prompt: Optional[Union[str, List[str]]] = None prompt: str | list[str] | None = None
negative_prompt: Optional[str] = None negative_prompt: str | None = None
prompt_path: Optional[str] = None prompt_path: str | None = None
output_path: str = "outputs/" output_path: str = "outputs/"
# Batch info # Batch info
@@ -53,7 +53,7 @@ class SamplingParam:
if self.prompt_path and not self.prompt_path.endswith(".txt"): if self.prompt_path and not self.prompt_path.endswith(".txt"):
raise ValueError("prompt_path must be a txt file") 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(): for key, value in source_dict.items():
if hasattr(self, key): if hasattr(self, key):
setattr(self, key, value) setattr(self, key, value)
+6 -6
View File
@@ -1,5 +1,6 @@
import os import os
from typing import Any, Callable, Dict, Optional from collections.abc import Callable
from typing import Any
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam, from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam) HunyuanSamplingParam)
@@ -14,7 +15,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
logger = init_logger(__name__) logger = init_logger(__name__)
# Registry maps specific model weights to their config classes # 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, "FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam, "hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam, "Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
@@ -26,7 +27,7 @@ SAMPLING_PARAM_REGISTRY: Dict[str, Any] = {
} }
# For determining pipeline type from model ID # 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(), "hunyuan": lambda id: "hunyuan" in id.lower(),
"wanpipeline": lambda id: "wanpipeline" in id.lower(), "wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(), "wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
@@ -35,7 +36,7 @@ SAMPLING_PARAM_DETECTOR: Dict[str, Callable[[str], bool]] = {
} }
# Fallback configs when exact match isn't found but architecture is detected # 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": "hunyuan":
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
"wanpipeline": "wanpipeline":
@@ -46,8 +47,7 @@ SAMPLING_FALLBACK_PARAM: Dict[str, Any] = {
} }
def get_sampling_param_cls_for_name( def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
pipeline_name_or_path: str) -> Optional[Any]:
"""Get the appropriate sampling param for specific pretrained weights.""" """Get the appropriate sampling param for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path): if os.path.exists(pipeline_name_or_path):
@@ -1,8 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # 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 # 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
import torch.distributed as dist import torch.distributed as dist
from torch.distributed import ProcessGroup from torch.distributed import ProcessGroup
@@ -18,8 +16,8 @@ class DeviceCommunicatorBase:
def __init__(self, def __init__(self,
cpu_group: ProcessGroup, cpu_group: ProcessGroup,
device: Optional[torch.device] = None, device: torch.device | None = None,
device_group: Optional[ProcessGroup] = None, device_group: ProcessGroup | None = None,
unique_name: str = ""): unique_name: str = ""):
self.device = device or torch.device("cpu") self.device = device or torch.device("cpu")
self.cpu_group = cpu_group self.cpu_group = cpu_group
@@ -66,7 +64,7 @@ class DeviceCommunicatorBase:
def gather(self, def gather(self,
input_: torch.Tensor, input_: torch.Tensor,
dst: int = 0, dst: int = 0,
dim: int = -1) -> Optional[torch.Tensor]: dim: int = -1) -> torch.Tensor | None:
""" """
NOTE: We assume that the input tensor is on the same device across NOTE: We assume that the input tensor is on the same device across
all the ranks. all the ranks.
@@ -170,7 +168,7 @@ class DeviceCommunicatorBase:
raise RuntimeError( raise RuntimeError(
"scatter_dim must be 1 or 2 and gather_dim must be 1 or 2") "scatter_dim must be 1 or 2 and gather_dim must be 1 or 2")
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None: def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way""" """Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank.""" """NOTE: `dst` is the local rank of the destination rank."""
if dst is None: if dst is None:
@@ -180,7 +178,7 @@ class DeviceCommunicatorBase:
def recv(self, def recv(self,
size: torch.Size, size: torch.Size,
dtype: torch.dtype, dtype: torch.dtype,
src: Optional[int] = None) -> torch.Tensor: src: int | None = None) -> torch.Tensor:
"""Receives a tensor from the source rank.""" """Receives a tensor from the source rank."""
"""NOTE: `src` is the local rank of the source rank.""" """NOTE: `src` is the local rank of the source rank."""
if src is None: if src is None:
@@ -1,8 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # 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 # 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 import torch
from torch.distributed import ProcessGroup from torch.distributed import ProcessGroup
@@ -14,15 +12,15 @@ class CudaCommunicator(DeviceCommunicatorBase):
def __init__(self, def __init__(self,
cpu_group: ProcessGroup, cpu_group: ProcessGroup,
device: Optional[torch.device] = None, device: torch.device | None = None,
device_group: Optional[ProcessGroup] = None, device_group: ProcessGroup | None = None,
unique_name: str = ""): unique_name: str = ""):
super().__init__(cpu_group, device, device_group, unique_name) super().__init__(cpu_group, device, device_group, unique_name)
from fastvideo.v1.distributed.device_communicators.pynccl import ( from fastvideo.v1.distributed.device_communicators.pynccl import (
PyNcclCommunicator) PyNcclCommunicator)
self.pynccl_comm: Optional[PyNcclCommunicator] = None self.pynccl_comm: PyNcclCommunicator | None = None
if self.world_size > 1: if self.world_size > 1:
self.pynccl_comm = PyNcclCommunicator( self.pynccl_comm = PyNcclCommunicator(
group=self.cpu_group, group=self.cpu_group,
@@ -42,7 +40,7 @@ class CudaCommunicator(DeviceCommunicatorBase):
torch.distributed.all_reduce(out, group=self.device_group) torch.distributed.all_reduce(out, group=self.device_group)
return out return out
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None: def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way""" """Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank.""" """NOTE: `dst` is the local rank of the destination rank."""
if dst is None: if dst is None:
@@ -57,7 +55,7 @@ class CudaCommunicator(DeviceCommunicatorBase):
def recv(self, def recv(self,
size: torch.Size, size: torch.Size,
dtype: torch.dtype, dtype: torch.dtype,
src: Optional[int] = None) -> torch.Tensor: src: int | None = None) -> torch.Tensor:
"""Receives a tensor from the source rank.""" """Receives a tensor from the source rank."""
"""NOTE: `src` is the local rank of the source rank.""" """NOTE: `src` is the local rank of the source rank."""
if src is None: if src is None:
@@ -1,8 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/pynccl.py # 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 region =====================
import torch import torch
import torch.distributed as dist import torch.distributed as dist
@@ -22,9 +20,9 @@ class PyNcclCommunicator:
def __init__( def __init__(
self, self,
group: Union[ProcessGroup, StatelessProcessGroup], group: ProcessGroup | StatelessProcessGroup,
device: Union[int, str, torch.device], device: int | str | torch.device,
library_path: Optional[str] = None, library_path: str | None = None,
): ):
""" """
Args: Args:
@@ -27,7 +27,7 @@
import ctypes import ctypes
import platform import platform
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Dict, List, Optional from typing import Any
import torch import torch
from torch.distributed import ReduceOp from torch.distributed import ReduceOp
@@ -124,7 +124,7 @@ class ncclRedOpTypeEnum:
class Function: class Function:
name: str name: str
restype: Any restype: Any
argtypes: List[Any] argtypes: list[Any]
class NCCLLibrary: class NCCLLibrary:
@@ -212,13 +212,13 @@ class NCCLLibrary:
# class attribute to store the mapping from the path to the library # class attribute to store the mapping from the path to the library
# to avoid loading the same library multiple times # 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 # class attribute to store the mapping from library path
# to the corresponding dictionary # to the corresponding dictionary
path_to_dict_mapping: Dict[str, Dict[str, Any]] = {} path_to_dict_mapping: dict[str, dict[str, Any]] = {}
def __init__(self, so_file: Optional[str] = None): def __init__(self, so_file: str | None = None):
so_file = so_file or find_nccl_library() so_file = so_file or find_nccl_library()
@@ -240,7 +240,7 @@ class NCCLLibrary:
raise e raise e
if so_file not in NCCLLibrary.path_to_dict_mapping: if so_file not in NCCLLibrary.path_to_dict_mapping:
_funcs: Dict[str, Any] = {} _funcs: dict[str, Any] = {}
for func in NCCLLibrary.exported_functions: for func in NCCLLibrary.exported_functions:
f = getattr(self.lib, func.name) f = getattr(self.lib, func.name)
f.restype = func.restype f.restype = func.restype
+45 -44
View File
@@ -27,10 +27,11 @@ import gc
import pickle import pickle
import weakref import weakref
from collections import namedtuple from collections import namedtuple
from collections.abc import Callable
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import dataclass from dataclasses import dataclass
from multiprocessing import shared_memory from multiprocessing import shared_memory
from typing import Any, Callable, Dict, List, Optional, Tuple, Union from typing import Any, Optional
from unittest.mock import patch from unittest.mock import patch
import torch import torch
@@ -57,15 +58,15 @@ TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
def _split_tensor_dict( def _split_tensor_dict(
tensor_dict: Dict[str, Union[torch.Tensor, Any]] tensor_dict: dict[str, torch.Tensor | Any]
) -> Tuple[List[Tuple[str, Any]], List[torch.Tensor]]: ) -> tuple[list[tuple[str, Any]], list[torch.Tensor]]:
"""Split the tensor dictionary into two parts: """Split the tensor dictionary into two parts:
1. A list of (key, value) pairs. If the value is a tensor, it is replaced 1. A list of (key, value) pairs. If the value is a tensor, it is replaced
by its metadata. by its metadata.
2. A list of tensors. 2. A list of tensors.
""" """
metadata_list: List[Tuple[str, Any]] = [] metadata_list: list[tuple[str, Any]] = []
tensor_list: List[torch.Tensor] = [] tensor_list: list[torch.Tensor] = []
for key, value in tensor_dict.items(): for key, value in tensor_dict.items():
if isinstance(value, torch.Tensor): if isinstance(value, torch.Tensor):
# Note: we cannot use `value.device` here, # Note: we cannot use `value.device` here,
@@ -81,7 +82,7 @@ def _split_tensor_dict(
return metadata_list, tensor_list return metadata_list, tensor_list
_group_name_counter: Dict[str, int] = {} _group_name_counter: dict[str, int] = {}
def _get_unique_name(name: str) -> str: def _get_unique_name(name: str) -> str:
@@ -97,7 +98,7 @@ def _get_unique_name(name: str) -> str:
return newname return newname
_groups: Dict[str, Callable[[], Optional["GroupCoordinator"]]] = {} _groups: dict[str, Callable[[], Optional["GroupCoordinator"]]] = {}
def _register_group(group: "GroupCoordinator") -> None: def _register_group(group: "GroupCoordinator") -> None:
@@ -128,7 +129,7 @@ class GroupCoordinator:
# available attributes: # available attributes:
rank: int # global rank 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 world_size: int # size of the group
# difference between `local_rank` and `rank_in_group`: # difference between `local_rank` and `rank_in_group`:
# if we have a group of size 4 across two nodes: # if we have a group of size 4 across two nodes:
@@ -143,16 +144,16 @@ class GroupCoordinator:
device_group: ProcessGroup # group for device communication device_group: ProcessGroup # group for device communication
use_device_communicator: bool # whether to use device communicator use_device_communicator: bool # whether to use device communicator
device_communicator: DeviceCommunicatorBase # device communicator device_communicator: DeviceCommunicatorBase # device communicator
mq_broadcaster: Optional[Any] # shared memory broadcaster mq_broadcaster: Any | None # shared memory broadcaster
def __init__( def __init__(
self, self,
group_ranks: List[List[int]], group_ranks: list[list[int]],
local_rank: int, local_rank: int,
torch_distributed_backend: Union[str, Backend], torch_distributed_backend: str | Backend,
use_device_communicator: bool, use_device_communicator: bool,
use_message_queue_broadcaster: bool = False, use_message_queue_broadcaster: bool = False,
group_name: Optional[str] = None, group_name: str | None = None,
): ):
group_name = group_name or "anonymous" group_name = group_name or "anonymous"
self.unique_name = _get_unique_name(group_name) self.unique_name = _get_unique_name(group_name)
@@ -243,8 +244,8 @@ class GroupCoordinator:
return self.ranks[(rank_in_group - 1) % world_size] return self.ranks[(rank_in_group - 1) % world_size]
@contextmanager @contextmanager
def graph_capture( def graph_capture(self,
self, graph_capture_context: Optional[GraphCaptureContext] = None): graph_capture_context: GraphCaptureContext | None = None):
if graph_capture_context is None: if graph_capture_context is None:
stream = torch.cuda.Stream() stream = torch.cuda.Stream()
graph_capture_context = GraphCaptureContext(stream) graph_capture_context = GraphCaptureContext(stream)
@@ -301,7 +302,7 @@ class GroupCoordinator:
def gather(self, def gather(self,
input_: torch.Tensor, input_: torch.Tensor,
dst: int = 0, dst: int = 0,
dim: int = -1) -> Optional[torch.Tensor]: dim: int = -1) -> torch.Tensor | None:
""" """
NOTE: We assume that the input tensor is on the same device across NOTE: We assume that the input tensor is on the same device across
all the ranks. all the ranks.
@@ -337,7 +338,7 @@ class GroupCoordinator:
group=self.device_group) group=self.device_group)
return input_ return input_
def broadcast_object(self, obj: Optional[Any] = None, src: int = 0): def broadcast_object(self, obj: Any | None = None, src: int = 0):
"""Broadcast the input object. """Broadcast the input object.
NOTE: `src` is the local rank of the source rank. NOTE: `src` is the local rank of the source rank.
""" """
@@ -362,9 +363,9 @@ class GroupCoordinator:
return recv[0] return recv[0]
def broadcast_object_list(self, def broadcast_object_list(self,
obj_list: List[Any], obj_list: list[Any],
src: int = 0, src: int = 0,
group: Optional[ProcessGroup] = None): group: ProcessGroup | None = None):
"""Broadcast the input object list. """Broadcast the input object list.
NOTE: `src` is the local rank of the source rank. NOTE: `src` is the local rank of the source rank.
""" """
@@ -444,11 +445,11 @@ class GroupCoordinator:
def broadcast_tensor_dict( def broadcast_tensor_dict(
self, self,
tensor_dict: Optional[Dict[str, Union[torch.Tensor, Any]]] = None, tensor_dict: dict[str, torch.Tensor | Any] | None = None,
src: int = 0, src: int = 0,
group: Optional[ProcessGroup] = None, group: ProcessGroup | None = None,
metadata_group: Optional[ProcessGroup] = None metadata_group: ProcessGroup | None = None
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]: ) -> dict[str, torch.Tensor | Any] | None:
"""Broadcast the input tensor dictionary. """Broadcast the input tensor dictionary.
NOTE: `src` is the local rank of the source rank. NOTE: `src` is the local rank of the source rank.
""" """
@@ -462,7 +463,7 @@ class GroupCoordinator:
rank_in_group = self.rank_in_group rank_in_group = self.rank_in_group
if rank_in_group == src: if rank_in_group == src:
metadata_list: List[Tuple[Any, Any]] = [] metadata_list: list[tuple[Any, Any]] = []
assert isinstance( assert isinstance(
tensor_dict, tensor_dict,
dict), (f"Expecting a dictionary, got {type(tensor_dict)}") dict), (f"Expecting a dictionary, got {type(tensor_dict)}")
@@ -529,10 +530,10 @@ class GroupCoordinator:
def send_tensor_dict( def send_tensor_dict(
self, self,
tensor_dict: Dict[str, Union[torch.Tensor, Any]], tensor_dict: dict[str, torch.Tensor | Any],
dst: Optional[int] = None, dst: int | None = None,
all_gather_group: Optional["GroupCoordinator"] = None, all_gather_group: Optional["GroupCoordinator"] = None,
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]: ) -> dict[str, torch.Tensor | Any] | None:
"""Send the input tensor dictionary. """Send the input tensor dictionary.
NOTE: `dst` is the local rank of the source rank. NOTE: `dst` is the local rank of the source rank.
""" """
@@ -552,7 +553,7 @@ class GroupCoordinator:
dst = (self.rank_in_group + 1) % self.world_size dst = (self.rank_in_group + 1) % self.world_size
assert dst < self.world_size, f"Invalid dst rank ({dst})" assert dst < self.world_size, f"Invalid dst rank ({dst})"
metadata_list: List[Tuple[Any, Any]] = [] metadata_list: list[tuple[Any, Any]] = []
assert isinstance( assert isinstance(
tensor_dict, tensor_dict,
dict), f"Expecting a dictionary, got {type(tensor_dict)}" dict), f"Expecting a dictionary, got {type(tensor_dict)}"
@@ -583,9 +584,9 @@ class GroupCoordinator:
def recv_tensor_dict( def recv_tensor_dict(
self, self,
src: Optional[int] = None, src: int | None = None,
all_gather_group: Optional["GroupCoordinator"] = None, all_gather_group: Optional["GroupCoordinator"] = None,
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]: ) -> dict[str, torch.Tensor | Any] | None:
"""Recv the input tensor dictionary. """Recv the input tensor dictionary.
NOTE: `src` is the local rank of the source rank. NOTE: `src` is the local rank of the source rank.
""" """
@@ -606,7 +607,7 @@ class GroupCoordinator:
assert src < self.world_size, f"Invalid src rank ({src})" assert src < self.world_size, f"Invalid src rank ({src})"
recv_metadata_list = self.recv_object(src=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: for key, value in recv_metadata_list:
if isinstance(value, TensorMetadata): if isinstance(value, TensorMetadata):
tensor = torch.empty(value.size, tensor = torch.empty(value.size,
@@ -656,7 +657,7 @@ class GroupCoordinator:
""" """
torch.distributed.barrier(group=self.cpu_group) torch.distributed.barrier(group=self.cpu_group)
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None: def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way""" """Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank.""" """NOTE: `dst` is the local rank of the destination rank."""
self.device_communicator.send(tensor, dst) self.device_communicator.send(tensor, dst)
@@ -664,7 +665,7 @@ class GroupCoordinator:
def recv(self, def recv(self,
size: torch.Size, size: torch.Size,
dtype: torch.dtype, dtype: torch.dtype,
src: Optional[int] = None) -> torch.Tensor: src: int | None = None) -> torch.Tensor:
"""Receives a tensor from the source rank.""" """Receives a tensor from the source rank."""
"""NOTE: `src` is the local rank of the source rank.""" """NOTE: `src` is the local rank of the source rank."""
return self.device_communicator.recv(size, dtype, src) return self.device_communicator.recv(size, dtype, src)
@@ -682,7 +683,7 @@ class GroupCoordinator:
self.mq_broadcaster = None self.mq_broadcaster = None
_WORLD: Optional[GroupCoordinator] = None _WORLD: GroupCoordinator | None = None
def get_world_group() -> GroupCoordinator: def get_world_group() -> GroupCoordinator:
@@ -690,7 +691,7 @@ def get_world_group() -> GroupCoordinator:
return _WORLD 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: backend: str) -> GroupCoordinator:
return GroupCoordinator( return GroupCoordinator(
group_ranks=[ranks], group_ranks=[ranks],
@@ -702,11 +703,11 @@ def init_world_group(ranks: List[int], local_rank: int,
def init_model_parallel_group( def init_model_parallel_group(
group_ranks: List[List[int]], group_ranks: list[list[int]],
local_rank: int, local_rank: int,
backend: str, backend: str,
use_message_queue_broadcaster: bool = False, use_message_queue_broadcaster: bool = False,
group_name: Optional[str] = None, group_name: str | None = None,
) -> GroupCoordinator: ) -> GroupCoordinator:
return GroupCoordinator( return GroupCoordinator(
@@ -719,7 +720,7 @@ def init_model_parallel_group(
) )
_TP: Optional[GroupCoordinator] = None _TP: GroupCoordinator | None = None
def get_tp_group() -> GroupCoordinator: def get_tp_group() -> GroupCoordinator:
@@ -778,7 +779,7 @@ def init_distributed_environment(
"world group already initialized with a different world size") "world group already initialized with a different world size")
_SP: Optional[GroupCoordinator] = None _SP: GroupCoordinator | None = None
def get_sp_group() -> GroupCoordinator: def get_sp_group() -> GroupCoordinator:
@@ -789,7 +790,7 @@ def get_sp_group() -> GroupCoordinator:
def initialize_model_parallel( def initialize_model_parallel(
tensor_model_parallel_size: int = 1, tensor_model_parallel_size: int = 1,
sequence_model_parallel_size: int = 1, sequence_model_parallel_size: int = 1,
backend: Optional[str] = None, backend: str | None = None,
) -> None: ) -> None:
""" """
Initialize model parallel groups. Initialize model parallel groups.
@@ -858,7 +859,7 @@ def get_sequence_model_parallel_rank() -> int:
def ensure_model_parallel_initialized( def ensure_model_parallel_initialized(
tensor_model_parallel_size: int, tensor_model_parallel_size: int,
sequence_model_parallel_size: int, sequence_model_parallel_size: int,
backend: Optional[str] = None, backend: str | None = None,
) -> None: ) -> None:
"""Helper to initialize model parallel groups if they are not initialized, """Helper to initialize model parallel groups if they are not initialized,
or ensure tensor-parallel, sequence-parallel sizes or ensure tensor-parallel, sequence-parallel sizes
@@ -969,8 +970,8 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
"torch._C._host_emptyCache() only available in Pytorch >=2.5") "torch._C._host_emptyCache() only available in Pytorch >=2.5")
def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup], def in_the_same_node_as(pg: ProcessGroup | StatelessProcessGroup,
source_rank: int = 0) -> List[bool]: source_rank: int = 0) -> list[bool]:
""" """
This is a collective operation that returns if each rank is in the same node This is a collective operation that returns if each rank is in the same node
as the source rank. It tests if processes are attached to the same as the source rank. It tests if processes are attached to the same
@@ -1056,7 +1057,7 @@ def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
def initialize_tensor_parallel_group( def initialize_tensor_parallel_group(
tensor_model_parallel_size: int = 1, tensor_model_parallel_size: int = 1,
backend: Optional[str] = None, backend: str | None = None,
group_name_suffix: str = "") -> GroupCoordinator: group_name_suffix: str = "") -> GroupCoordinator:
"""Initialize a tensor parallel group for a specific model. """Initialize a tensor parallel group for a specific model.
@@ -1120,7 +1121,7 @@ def initialize_tensor_parallel_group(
def initialize_sequence_parallel_group( def initialize_sequence_parallel_group(
sequence_model_parallel_size: int = 1, sequence_model_parallel_size: int = 1,
backend: Optional[str] = None, backend: str | None = None,
group_name_suffix: str = "") -> GroupCoordinator: group_name_suffix: str = "") -> GroupCoordinator:
"""Initialize a sequence parallel group for a specific model. """Initialize a sequence parallel group for a specific model.
+7 -6
View File
@@ -9,7 +9,8 @@ import dataclasses
import pickle import pickle
import time import time
from collections import deque from collections import deque
from typing import Any, Deque, Dict, Optional, Sequence, Tuple from collections.abc import Sequence
from typing import Any
import torch import torch
from torch.distributed import TCPStore from torch.distributed import TCPStore
@@ -72,15 +73,15 @@ class StatelessProcessGroup:
data_expiration_seconds: int = 3600 # 1 hour data_expiration_seconds: int = 3600 # 1 hour
# dst rank -> counter # 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 # 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_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) default_factory=dict)
# A deque to store the data entries, with key and timestamp. # 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): def __post_init__(self):
assert self.rank < self.world_size assert self.rank < self.world_size
@@ -114,7 +115,7 @@ class StatelessProcessGroup:
self.recv_src_counter[src] += 1 self.recv_src_counter[src] += 1
return obj return obj
def broadcast_obj(self, obj: Optional[Any], src: int) -> Any: def broadcast_obj(self, obj: Any | None, src: int) -> Any:
"""Broadcast an object from a source rank to all other ranks. """Broadcast an object from a source rank to all other ranks.
It does not clean up after all ranks have received the object. It does not clean up after all ranks have received the object.
Use it for limited times, e.g., for initialization. Use it for limited times, e.g., for initialization.
+6 -6
View File
@@ -4,7 +4,7 @@
import argparse import argparse
import dataclasses import dataclasses
import os import os
from typing import Any, Dict, List, Optional, cast from typing import Any, cast
from fastvideo import PipelineConfig, VideoGenerator from fastvideo import PipelineConfig, VideoGenerator
from fastvideo.v1.configs.sample.base import SamplingParam 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.init_arg_names = self._get_init_arg_names()
self.generation_arg_names = self._get_generation_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""" """Get names of arguments for VideoGenerator initialization"""
return ["num_gpus", "tp_size", "sp_size", "model_path"] 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""" """Get names of arguments for generate_video method"""
return [field.name for field in dataclasses.fields(SamplingParam)] return [field.name for field in dataclasses.fields(SamplingParam)]
@@ -130,13 +130,13 @@ class GenerateSubcommand(CLISubcommand):
return cast(FlexibleArgumentParser, generate_parser) return cast(FlexibleArgumentParser, generate_parser)
def cmd_init() -> List[CLISubcommand]: def cmd_init() -> list[CLISubcommand]:
return [GenerateSubcommand()] return [GenerateSubcommand()]
def update_config_from_args(config: Any, def update_config_from_args(config: Any,
args_dict: Dict[str, Any], args_dict: dict[str, Any],
prefix: Optional[str] = None) -> None: prefix: str | None = None) -> None:
""" """
Update configuration object from arguments dictionary. Update configuration object from arguments dictionary.
+1 -3
View File
@@ -1,14 +1,12 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/main.py # 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.cli_types import CLISubcommand
from fastvideo.v1.entrypoints.cli.generate import cmd_init as generate_cmd_init from fastvideo.v1.entrypoints.cli.generate import cmd_init as generate_cmd_init
from fastvideo.v1.utils import FlexibleArgumentParser from fastvideo.v1.utils import FlexibleArgumentParser
def cmd_init() -> List[CLISubcommand]: def cmd_init() -> list[CLISubcommand]:
"""Initialize all commands from separate modules""" """Initialize all commands from separate modules"""
commands = [] commands = []
commands.extend(generate_cmd_init()) commands.extend(generate_cmd_init())
+2 -3
View File
@@ -4,7 +4,6 @@ import argparse
import os import os
import subprocess import subprocess
import sys import sys
from typing import List, Optional
from fastvideo.v1.logger import init_logger from fastvideo.v1.logger import init_logger
@@ -19,8 +18,8 @@ class RaiseNotImplementedAction(argparse.Action):
def launch_distributed(num_gpus: int, def launch_distributed(num_gpus: int,
args: List[str], args: list[str],
master_port: Optional[int] = None) -> int: master_port: int | None = None) -> int:
""" """
Launch a distributed job with the given arguments Launch a distributed job with the given arguments
+6 -8
View File
@@ -10,7 +10,7 @@ import gc
import math import math
import os import os
import time import time
from typing import Any, Dict, List, Optional, Union from typing import Any
import imageio import imageio
import numpy as np import numpy as np
@@ -53,11 +53,9 @@ class VideoGenerator:
@classmethod @classmethod
def from_pretrained(cls, def from_pretrained(cls,
model_path: str, model_path: str,
device: Optional[str] = None, device: str | None = None,
torch_dtype: Optional[torch.dtype] = None, torch_dtype: torch.dtype | None = None,
pipeline_config: Optional[ pipeline_config: str | PipelineConfig | None = None,
Union[str
| PipelineConfig]] = None,
**kwargs) -> "VideoGenerator": **kwargs) -> "VideoGenerator":
""" """
Create a video generator from a pretrained model. Create a video generator from a pretrained model.
@@ -128,9 +126,9 @@ class VideoGenerator:
def generate_video( def generate_video(
self, self,
prompt: str, prompt: str,
sampling_param: Optional[SamplingParam] = None, sampling_param: SamplingParam | None = None,
**kwargs, **kwargs,
) -> Union[Dict[str, Any], List[np.ndarray]]: ) -> dict[str, Any] | list[np.ndarray]:
""" """
Generate a video based on the given prompt. Generate a video based on the given prompt.
+13 -12
View File
@@ -2,28 +2,29 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/envs.py # Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/envs.py
import os import os
from typing import TYPE_CHECKING, Any, Callable, Dict, Optional from collections.abc import Callable
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING: if TYPE_CHECKING:
FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60 FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60
FASTVIDEO_NCCL_SO_PATH: Optional[str] = None FASTVIDEO_NCCL_SO_PATH: str | None = None
LD_LIBRARY_PATH: Optional[str] = None LD_LIBRARY_PATH: str | None = None
LOCAL_RANK: int = 0 LOCAL_RANK: int = 0
CUDA_VISIBLE_DEVICES: Optional[str] = None CUDA_VISIBLE_DEVICES: str | None = None
FASTVIDEO_CACHE_ROOT: str = os.path.expanduser("~/.cache/fastvideo") FASTVIDEO_CACHE_ROOT: str = os.path.expanduser("~/.cache/fastvideo")
FASTVIDEO_CONFIG_ROOT: str = os.path.expanduser("~/.config/fastvideo") FASTVIDEO_CONFIG_ROOT: str = os.path.expanduser("~/.config/fastvideo")
FASTVIDEO_CONFIGURE_LOGGING: int = 1 FASTVIDEO_CONFIGURE_LOGGING: int = 1
FASTVIDEO_LOGGING_LEVEL: str = "INFO" FASTVIDEO_LOGGING_LEVEL: str = "INFO"
FASTVIDEO_LOGGING_PREFIX: str = "" FASTVIDEO_LOGGING_PREFIX: str = ""
FASTVIDEO_LOGGING_CONFIG_PATH: Optional[str] = None FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
FASTVIDEO_TRACE_FUNCTION: int = 0 FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: Optional[str] = None FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_ATTENTION_CONFIG: Optional[str] = None FASTVIDEO_ATTENTION_CONFIG: str | None = None
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork" FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
FASTVIDEO_TARGET_DEVICE: str = "cuda" FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: Optional[str] = None MAX_JOBS: str | None = None
NVCC_THREADS: Optional[str] = None NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: Optional[str] = None CMAKE_BUILD_TYPE: str | None = None
VERBOSE: bool = False VERBOSE: bool = False
FASTVIDEO_SERVER_DEV_MODE: bool = False FASTVIDEO_SERVER_DEV_MODE: bool = False
@@ -42,7 +43,7 @@ def get_default_config_root() -> str:
) )
def maybe_convert_int(value: Optional[str]) -> Optional[int]: def maybe_convert_int(value: str | None) -> int | None:
if value is None: if value is None:
return None return None
return int(value) return int(value)
@@ -53,7 +54,7 @@ def maybe_convert_int(value: Optional[str]) -> Optional[int]:
# begin-env-vars-definition # begin-env-vars-definition
environment_variables: Dict[str, Callable[[], Any]] = { environment_variables: dict[str, Callable[[], Any]] = {
# ================== Installation Time Env Vars ================== # ================== Installation Time Env Vars ==================
+17 -16
View File
@@ -4,9 +4,10 @@
import argparse import argparse
import dataclasses import dataclasses
from collections.abc import Callable
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import field from dataclasses import field
from typing import Any, Callable, List, Optional, Tuple from typing import Any
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.logger import init_logger from fastvideo.v1.logger import init_logger
@@ -38,17 +39,17 @@ class FastVideoArgs:
# HuggingFace specific parameters # HuggingFace specific parameters
trust_remote_code: bool = False trust_remote_code: bool = False
revision: Optional[str] = None revision: str | None = None
# Parallelism # Parallelism
num_gpus: int = 1 num_gpus: int = 1
tp_size: Optional[int] = None tp_size: int | None = None
sp_size: Optional[int] = None sp_size: int | None = None
dist_timeout: Optional[int] = None # timeout for torch.distributed dist_timeout: int | None = None # timeout for torch.distributed
# Video generation parameters # Video generation parameters
embedded_cfg_scale: float = 6.0 embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None flow_shift: float | None = None
output_type: str = "pil" output_type: str = "pil"
@@ -72,32 +73,32 @@ class FastVideoArgs:
"fp16", "fp16",
"fp16", "fp16",
) )
text_encoder_precisions: Tuple[str, ...] = field( text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS) default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
text_encoder_configs: Tuple[EncoderConfig, ...] = field( text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), )) 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, )) 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, )) default_factory=lambda: (postprocess_text, ))
# STA (Spatial-Temporal Attention) parameters # STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None mask_strategy_file_path: str | None = None
enable_torch_compile: bool = False enable_torch_compile: bool = False
use_cpu_offload: bool = False use_cpu_offload: bool = False
disable_autocast: bool = False disable_autocast: bool = False
# StepVideo specific parameters # StepVideo specific parameters
pos_magic: Optional[str] = None pos_magic: str | None = None
neg_magic: Optional[str] = None neg_magic: str | None = None
timesteps_scale: Optional[bool] = None timesteps_scale: bool | None = None
# Logging # Logging
log_level: str = "info" log_level: str = "info"
# Inference parameters # Inference parameters
device_str: Optional[str] = None device_str: str | None = None
device = None device = None
def __post_init__(self): def __post_init__(self):
@@ -376,7 +377,7 @@ class FastVideoArgs:
_current_fastvideo_args = None _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. Prepare the inference arguments from the command line arguments.
+5 -5
View File
@@ -5,7 +5,7 @@ import time
from collections import defaultdict from collections import defaultdict
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import dataclass from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING
import torch import torch
@@ -37,10 +37,10 @@ class ForwardContext:
# attn_layers: Dict[str, Any] # attn_layers: Dict[str, Any]
# TODO: extend to support per-layer dynamic forward context # TODO: extend to support per-layer dynamic forward context
attn_metadata: "AttentionMetadata" # set dynamically for each forward pass attn_metadata: "AttentionMetadata" # set dynamically for each forward pass
forward_batch: Optional[ForwardBatch] = None forward_batch: ForwardBatch | None = None
_forward_context: Optional[ForwardContext] = None _forward_context: ForwardContext | None = None
def get_forward_context() -> ForwardContext: def get_forward_context() -> ForwardContext:
@@ -55,8 +55,8 @@ def get_forward_context() -> ForwardContext:
@contextmanager @contextmanager
def set_forward_context(current_timestep, def set_forward_context(current_timestep,
attn_metadata, attn_metadata,
forward_batch: Optional[ForwardBatch] = None, forward_batch: ForwardBatch | None = None,
fastvideo_args: Optional[FastVideoArgs] = None): fastvideo_args: FastVideoArgs | None = None):
"""A context manager that stores the current forward context, """A context manager that stores the current forward context,
can be attention metadata, etc. can be attention metadata, etc.
Here we can inject common logic for every model forward pass. 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 import time
from typing import Any, Dict from typing import Any
import torch import torch
@@ -83,7 +83,7 @@ class InferenceEngine:
self, self,
prompt: str, prompt: str,
fastvideo_args: FastVideoArgs, fastvideo_args: FastVideoArgs,
) -> Dict[str, Any]: ) -> dict[str, Any]:
""" """
Run inference with the pipeline. Run inference with the pipeline.
@@ -96,7 +96,7 @@ class InferenceEngine:
Returns: Returns:
A dictionary containing the generated videos and metadata. 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 num_videos_per_prompt = fastvideo_args.num_videos
seed = fastvideo_args.seed seed = fastvideo_args.seed
+3 -2
View File
@@ -1,7 +1,8 @@
# SPDX-License-Identifier: Apache-2.0 # 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 # Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/custom_op.py
from typing import Any, Callable, Dict, Type from collections.abc import Callable
from typing import Any
import torch.nn as nn import torch.nn as nn
@@ -81,7 +82,7 @@ class CustomOp(nn.Module):
# Examples: # Examples:
# - MyOp.enabled() # - MyOp.enabled()
# - op_registry["my_op"].enabled() # - op_registry["my_op"].enabled()
op_registry: Dict[str, Type['CustomOp']] = {} op_registry: dict[str, type['CustomOp']] = {}
# Decorator to register custom ops. # Decorator to register custom ops.
@classmethod @classmethod
+4 -5
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # 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 # Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/layernorm.py
"""Custom normalization layers.""" """Custom normalization layers."""
from typing import Optional, Tuple, Union
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -22,7 +21,7 @@ class RMSNorm(CustomOp):
hidden_size: int, hidden_size: int,
eps: float = 1e-6, eps: float = 1e-6,
dtype: torch.dtype = torch.float32, dtype: torch.dtype = torch.float32,
var_hidden_size: Optional[int] = None, var_hidden_size: int | None = None,
has_weight: bool = True, has_weight: bool = True,
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -40,8 +39,8 @@ class RMSNorm(CustomOp):
def forward_native( def forward_native(
self, self,
x: torch.Tensor, x: torch.Tensor,
residual: Optional[torch.Tensor] = None, residual: torch.Tensor | None = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""PyTorch-native implementation equivalent to forward().""" """PyTorch-native implementation equivalent to forward()."""
orig_dtype = x.dtype orig_dtype = x.dtype
x = x.to(torch.float32) x = x.to(torch.float32)
@@ -130,7 +129,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
def forward(self, residual: torch.Tensor, x: torch.Tensor, def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor, shift: 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 Apply gated residual connection, followed by layernorm and
scale/shift in a single fused operation. scale/shift in a single fused operation.
+34 -42
View File
@@ -2,7 +2,6 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/linear.py # Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/linear.py
from abc import abstractmethod from abc import abstractmethod
from typing import Optional, Union
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
@@ -40,7 +39,7 @@ WEIGHT_LOADER_V2_SUPPORTED = [
def adjust_scalar_to_fused_array( def adjust_scalar_to_fused_array(
param: torch.Tensor, loaded_weight: torch.Tensor, param: torch.Tensor, loaded_weight: torch.Tensor,
shard_id: Union[str, int]) -> tuple[torch.Tensor, torch.Tensor]: shard_id: str | int) -> tuple[torch.Tensor, torch.Tensor]:
"""For fused modules (QKV and MLP) we have an array of length """For fused modules (QKV and MLP) we have an array of length
N that holds 1 scale for each "logical" matrix. So the param N that holds 1 scale for each "logical" matrix. So the param
is an array of length N. The loaded_weight corresponds to is an array of length N. The loaded_weight corresponds to
@@ -91,7 +90,7 @@ class LinearMethodBase(QuantizeMethodBase):
def apply(self, def apply(self,
layer: torch.nn.Module, layer: torch.nn.Module,
x: torch.Tensor, x: torch.Tensor,
bias: Optional[torch.Tensor] = None) -> torch.Tensor: bias: torch.Tensor | None = None) -> torch.Tensor:
"""Apply the weights in layer to the input tensor. """Apply the weights in layer to the input tensor.
Expects create_weights to have been called before on the layer.""" Expects create_weights to have been called before on the layer."""
raise NotImplementedError raise NotImplementedError
@@ -116,7 +115,7 @@ class UnquantizedLinearMethod(LinearMethodBase):
def apply(self, def apply(self,
layer: torch.nn.Module, layer: torch.nn.Module,
x: torch.Tensor, x: torch.Tensor,
bias: Optional[torch.Tensor] = None) -> torch.Tensor: bias: torch.Tensor | None = None) -> torch.Tensor:
return F.linear(x, layer.weight, bias) return F.linear(x, layer.weight, bias)
@@ -138,8 +137,8 @@ class LinearBase(torch.nn.Module):
input_size: int, input_size: int,
output_size: int, output_size: int,
skip_bias_add: bool = False, skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None, params_dtype: torch.dtype | None = None,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
@@ -152,14 +151,13 @@ class LinearBase(torch.nn.Module):
params_dtype = torch.get_default_dtype() params_dtype = torch.get_default_dtype()
self.params_dtype = params_dtype self.params_dtype = params_dtype
if quant_config is None: if quant_config is None:
self.quant_method: Optional[ self.quant_method: QuantizeMethodBase | None = UnquantizedLinearMethod(
QuantizeMethodBase] = UnquantizedLinearMethod() )
else: else:
self.quant_method = quant_config.get_quant_method(self, self.quant_method = quant_config.get_quant_method(self,
prefix=prefix) prefix=prefix)
def forward(self, def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
x: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
raise NotImplementedError raise NotImplementedError
@@ -182,8 +180,8 @@ class ReplicatedLinear(LinearBase):
output_size: int, output_size: int,
bias: bool = True, bias: bool = True,
skip_bias_add: bool = False, skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None, params_dtype: torch.dtype | None = None,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = ""): prefix: str = ""):
super().__init__(input_size, super().__init__(input_size,
output_size, output_size,
@@ -223,8 +221,7 @@ class ReplicatedLinear(LinearBase):
f"to a parameter of size {param.size()}") f"to a parameter of size {param.size()}")
param.data.copy_(loaded_weight) param.data.copy_(loaded_weight)
def forward(self, def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
x: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
bias = self.bias if not self.skip_bias_add else None bias = self.bias if not self.skip_bias_add else None
assert self.quant_method is not None assert self.quant_method is not None
output = self.quant_method.apply(self, x, bias) output = self.quant_method.apply(self, x, bias)
@@ -268,9 +265,9 @@ class ColumnParallelLinear(LinearBase):
bias: bool = True, bias: bool = True,
gather_output: bool = False, gather_output: bool = False,
skip_bias_add: bool = False, skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None, params_dtype: torch.dtype | None = None,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
output_sizes: Optional[list[int]] = None, output_sizes: list[int] | None = None,
prefix: str = ""): prefix: str = ""):
# Divide the weight matrix along the last dimension. # Divide the weight matrix along the last dimension.
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_tensor_model_parallel_world_size()
@@ -345,9 +342,8 @@ class ColumnParallelLinear(LinearBase):
loaded_weight = loaded_weight.reshape(1) loaded_weight = loaded_weight.reshape(1)
param.load_column_parallel_weight(loaded_weight=loaded_weight) param.load_column_parallel_weight(loaded_weight=loaded_weight)
def forward( def forward(self,
self, input_: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
input_: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
bias = self.bias if not self.skip_bias_add else None bias = self.bias if not self.skip_bias_add else None
# Matrix multiply. # Matrix multiply.
@@ -399,8 +395,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
bias: bool = True, bias: bool = True,
gather_output: bool = False, gather_output: bool = False,
skip_bias_add: bool = False, skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None, params_dtype: torch.dtype | None = None,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = ""): prefix: str = ""):
self.output_sizes = output_sizes self.output_sizes = output_sizes
tp_size = get_tensor_model_parallel_world_size() tp_size = get_tensor_model_parallel_world_size()
@@ -417,7 +413,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
def weight_loader(self, def weight_loader(self,
param: Parameter, param: Parameter,
loaded_weight: torch.Tensor, loaded_weight: torch.Tensor,
loaded_shard_id: Optional[int] = None) -> None: loaded_shard_id: int | None = None) -> None:
param_data = param.data param_data = param.data
output_dim = getattr(param, "output_dim", None) output_dim = getattr(param, "output_dim", None)
@@ -510,10 +506,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
# Special case for Quantization. # Special case for Quantization.
# If quantized, we need to adjust the offset and size to account # If quantized, we need to adjust the offset and size to account
# for the packing. # for the packing.
if isinstance( if isinstance(param, PackedColumnParameter | PackedvLLMParameter
param, ) and param.packed_dim == param.output_dim:
(PackedColumnParameter,
PackedvLLMParameter)) and param.packed_dim == param.output_dim:
shard_size, shard_offset = \ shard_size, shard_offset = \
param.adjust_shard_indexes_for_packing( param.adjust_shard_indexes_for_packing(
shard_size=shard_size, shard_offset=shard_offset) shard_size=shard_size, shard_offset=shard_offset)
@@ -525,7 +519,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
def weight_loader_v2(self, def weight_loader_v2(self,
param: BasevLLMParameter, param: BasevLLMParameter,
loaded_weight: torch.Tensor, loaded_weight: torch.Tensor,
loaded_shard_id: Optional[int] = None) -> None: loaded_shard_id: int | None = None) -> None:
if loaded_shard_id is None: if loaded_shard_id is None:
if isinstance(param, PerTensorScaleParameter): if isinstance(param, PerTensorScaleParameter):
param.load_merged_column_weight(loaded_weight=loaded_weight, param.load_merged_column_weight(loaded_weight=loaded_weight,
@@ -598,11 +592,11 @@ class QKVParallelLinear(ColumnParallelLinear):
hidden_size: int, hidden_size: int,
head_size: int, head_size: int,
total_num_heads: int, total_num_heads: int,
total_num_kv_heads: Optional[int] = None, total_num_kv_heads: int | None = None,
bias: bool = True, bias: bool = True,
skip_bias_add: bool = False, skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None, params_dtype: torch.dtype | None = None,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = ""): prefix: str = ""):
self.hidden_size = hidden_size self.hidden_size = hidden_size
self.head_size = head_size self.head_size = head_size
@@ -637,7 +631,7 @@ class QKVParallelLinear(ColumnParallelLinear):
quant_config=quant_config, quant_config=quant_config,
prefix=prefix) prefix=prefix)
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> Optional[int]: def _get_shard_offset_mapping(self, loaded_shard_id: str) -> int | None:
shard_offset_mapping = { shard_offset_mapping = {
"q": 0, "q": 0,
"k": self.num_heads * self.head_size, "k": self.num_heads * self.head_size,
@@ -646,7 +640,7 @@ class QKVParallelLinear(ColumnParallelLinear):
} }
return shard_offset_mapping.get(loaded_shard_id) return shard_offset_mapping.get(loaded_shard_id)
def _get_shard_size_mapping(self, loaded_shard_id: str) -> Optional[int]: def _get_shard_size_mapping(self, loaded_shard_id: str) -> int | None:
shard_size_mapping = { shard_size_mapping = {
"q": self.num_heads * self.head_size, "q": self.num_heads * self.head_size,
"k": self.num_kv_heads * self.head_size, "k": self.num_kv_heads * self.head_size,
@@ -679,10 +673,8 @@ class QKVParallelLinear(ColumnParallelLinear):
# Special case for Quantization. # Special case for Quantization.
# If quantized, we need to adjust the offset and size to account # If quantized, we need to adjust the offset and size to account
# for the packing. # for the packing.
if isinstance( if isinstance(param, PackedColumnParameter | PackedvLLMParameter
param, ) and param.packed_dim == param.output_dim:
(PackedColumnParameter,
PackedvLLMParameter)) and param.packed_dim == param.output_dim:
shard_size, shard_offset = \ shard_size, shard_offset = \
param.adjust_shard_indexes_for_packing( param.adjust_shard_indexes_for_packing(
shard_size=shard_size, shard_offset=shard_offset) shard_size=shard_size, shard_offset=shard_offset)
@@ -694,7 +686,7 @@ class QKVParallelLinear(ColumnParallelLinear):
def weight_loader_v2(self, def weight_loader_v2(self,
param: BasevLLMParameter, param: BasevLLMParameter,
loaded_weight: torch.Tensor, loaded_weight: torch.Tensor,
loaded_shard_id: Optional[str] = None): loaded_shard_id: str | None = None):
if loaded_shard_id is None: # special case for certain models if loaded_shard_id is None: # special case for certain models
if isinstance(param, PerTensorScaleParameter): if isinstance(param, PerTensorScaleParameter):
param.load_qkv_weight(loaded_weight=loaded_weight, shard_id=0) param.load_qkv_weight(loaded_weight=loaded_weight, shard_id=0)
@@ -720,7 +712,7 @@ class QKVParallelLinear(ColumnParallelLinear):
def weight_loader(self, def weight_loader(self,
param: Parameter, param: Parameter,
loaded_weight: torch.Tensor, loaded_weight: torch.Tensor,
loaded_shard_id: Optional[str] = None): loaded_shard_id: str | None = None):
param_data = param.data param_data = param.data
output_dim = getattr(param, "output_dim", None) output_dim = getattr(param, "output_dim", None)
@@ -845,9 +837,9 @@ class RowParallelLinear(LinearBase):
bias: bool = True, bias: bool = True,
input_is_parallel: bool = True, input_is_parallel: bool = True,
skip_bias_add: bool = False, skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None, params_dtype: torch.dtype | None = None,
reduce_results: bool = True, reduce_results: bool = True,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = ""): prefix: str = ""):
# Divide the weight matrix along the first dimension. # Divide the weight matrix along the first dimension.
self.tp_rank = get_tensor_model_parallel_rank() self.tp_rank = get_tensor_model_parallel_rank()
@@ -921,7 +913,7 @@ class RowParallelLinear(LinearBase):
param.load_row_parallel_weight(loaded_weight=loaded_weight) param.load_row_parallel_weight(loaded_weight=loaded_weight)
def forward(self, input_) -> tuple[torch.Tensor, Optional[Parameter]]: def forward(self, input_) -> tuple[torch.Tensor, Parameter | None]:
if self.input_is_parallel: if self.input_is_parallel:
input_parallel = input_ input_parallel = input_
else: else:
+2 -4
View File
@@ -1,7 +1,5 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from typing import Optional
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -18,10 +16,10 @@ class MLP(nn.Module):
self, self,
input_dim: int, input_dim: int,
mlp_hidden_dim: int, mlp_hidden_dim: int,
output_dim: Optional[int] = None, output_dim: int | None = None,
bias: bool = True, bias: bool = True,
act_type: str = "gelu_pytorch_tanh", act_type: str = "gelu_pytorch_tanh",
dtype: Optional[torch.dtype] = None, dtype: torch.dtype | None = None,
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
@@ -3,7 +3,7 @@
import inspect import inspect
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Optional from typing import TYPE_CHECKING, Any
import torch import torch
from torch import nn from torch import nn
@@ -105,8 +105,8 @@ class QuantizationConfig(ABC):
raise NotImplementedError raise NotImplementedError
@classmethod @classmethod
def override_quantization_method( def override_quantization_method(cls, hf_quant_cfg,
cls, hf_quant_cfg, user_quant) -> Optional[QuantizationMethods]: user_quant) -> QuantizationMethods | None:
""" """
Detects if this quantization method can support a given checkpoint Detects if this quantization method can support a given checkpoint
format by overriding the user specified quantization method -- format by overriding the user specified quantization method --
@@ -135,7 +135,7 @@ class QuantizationConfig(ABC):
@abstractmethod @abstractmethod
def get_quant_method(self, layer: torch.nn.Module, def get_quant_method(self, layer: torch.nn.Module,
prefix: str) -> Optional[QuantizeMethodBase]: prefix: str) -> QuantizeMethodBase | None:
"""Get the quantize method to use for the quantized layer. """Get the quantize method to use for the quantized layer.
Args: Args:
@@ -147,5 +147,5 @@ class QuantizationConfig(ABC):
""" """
raise NotImplementedError raise NotImplementedError
def get_cache_scale(self, name: str) -> Optional[str]: def get_cache_scale(self, name: str) -> str | None:
return None return None
+20 -20
View File
@@ -23,7 +23,7 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
"""Rotary Positional Embeddings.""" """Rotary Positional Embeddings."""
from typing import Any, Dict, List, Optional, Tuple, Union from typing import Any
import torch import torch
@@ -84,7 +84,7 @@ class RotaryEmbedding(CustomOp):
head_size: int, head_size: int,
rotary_dim: int, rotary_dim: int,
max_position_embeddings: int, max_position_embeddings: int,
base: Union[int, float], base: int | float,
is_neox_style: bool, is_neox_style: bool,
dtype: torch.dtype, dtype: torch.dtype,
) -> None: ) -> None:
@@ -101,7 +101,7 @@ class RotaryEmbedding(CustomOp):
self.cos_sin_cache: torch.Tensor self.cos_sin_cache: torch.Tensor
self.register_buffer("cos_sin_cache", cache, persistent=False) self.register_buffer("cos_sin_cache", cache, persistent=False)
def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor: def _compute_inv_freq(self, base: int | float) -> torch.Tensor:
"""Compute the inverse frequency.""" """Compute the inverse frequency."""
# NOTE(woosuk): To exactly match the HF implementation, we need to # 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 # use CPU to compute the cache and then move it to GPU. However, we
@@ -127,8 +127,8 @@ class RotaryEmbedding(CustomOp):
positions: torch.Tensor, positions: torch.Tensor,
query: torch.Tensor, query: torch.Tensor,
key: torch.Tensor, key: torch.Tensor,
offsets: Optional[torch.Tensor] = None, offsets: torch.Tensor | None = None,
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
"""A PyTorch-native implementation of forward().""" """A PyTorch-native implementation of forward()."""
if offsets is not None: if offsets is not None:
positions = positions + offsets positions = positions + offsets
@@ -159,7 +159,7 @@ class RotaryEmbedding(CustomOp):
return s return s
def _to_tuple(x: Union[int, Tuple[int, ...]], dim: int = 2) -> Tuple[int, ...]: def _to_tuple(x: int | tuple[int, ...], dim: int = 2) -> tuple[int, ...]:
if isinstance(x, int): if isinstance(x, int):
return (x, ) * dim return (x, ) * dim
elif len(x) == dim: elif len(x) == dim:
@@ -168,8 +168,8 @@ def _to_tuple(x: Union[int, Tuple[int, ...]], dim: int = 2) -> Tuple[int, ...]:
raise ValueError(f"Expected length {dim} or int, but got {x}") raise ValueError(f"Expected length {dim} or int, but got {x}")
def get_meshgrid_nd(start: Union[int, Tuple[int, ...]], def get_meshgrid_nd(start: int | tuple[int, ...],
*args: Union[int, Tuple[int, ...]], *args: int | tuple[int, ...],
dim: int = 2) -> torch.Tensor: dim: int = 2) -> torch.Tensor:
""" """
Get n-D meshgrid with start, stop and num. Get n-D meshgrid with start, stop and num.
@@ -217,12 +217,12 @@ def get_meshgrid_nd(start: Union[int, Tuple[int, ...]],
def get_1d_rotary_pos_embed( def get_1d_rotary_pos_embed(
dim: int, dim: int,
pos: Union[torch.FloatTensor, int], pos: torch.FloatTensor | int,
theta: float = 10000.0, theta: float = 10000.0,
theta_rescale_factor: float = 1.0, theta_rescale_factor: float = 1.0,
interpolation_factor: float = 1.0, interpolation_factor: float = 1.0,
dtype: torch.dtype = torch.float32, 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. Precompute the frequency tensor for complex exponential (cis) with given dimensions.
(Note: `cis` means `cos + i * sin`, where i is the imaginary unit.) (Note: `cis` means `cos + i * sin`, where i is the imaginary unit.)
@@ -261,13 +261,13 @@ def get_nd_rotary_pos_embed(
start, start,
*args, *args,
theta=10000.0, theta=10000.0,
theta_rescale_factor: Union[float, List[float]] = 1.0, theta_rescale_factor: float | list[float] = 1.0,
interpolation_factor: Union[float, List[float]] = 1.0, interpolation_factor: float | list[float] = 1.0,
shard_dim: int = 0, shard_dim: int = 0,
sp_rank: int = 0, sp_rank: int = 0,
sp_world_size: int = 1, sp_world_size: int = 1,
dtype: torch.dtype = torch.float32, 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. 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. Supports sequence parallelism by allowing sharding of a specific dimension.
@@ -324,7 +324,7 @@ def get_nd_rotary_pos_embed(
else: else:
grid = full_grid 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) theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
elif isinstance(theta_rescale_factor, elif isinstance(theta_rescale_factor,
list) and len(theta_rescale_factor) == 1: list) and len(theta_rescale_factor) == 1:
@@ -333,7 +333,7 @@ def get_nd_rotary_pos_embed(
rope_dim_list rope_dim_list
), "len(theta_rescale_factor) should equal to len(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) interpolation_factor = [interpolation_factor] * len(rope_dim_list)
elif isinstance(interpolation_factor, elif isinstance(interpolation_factor,
list) and len(interpolation_factor) == 1: list) and len(interpolation_factor) == 1:
@@ -370,7 +370,7 @@ def get_rotary_pos_embed(
interpolation_factor=1.0, interpolation_factor=1.0,
shard_dim: int = 0, shard_dim: int = 0,
dtype: torch.dtype = torch.float32, dtype: torch.dtype = torch.float32,
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
""" """
Generate rotary positional embeddings for the given sizes. Generate rotary positional embeddings for the given sizes.
@@ -417,17 +417,17 @@ def get_rotary_pos_embed(
return freqs_cos, freqs_sin return freqs_cos, freqs_sin
_ROPE_DICT: Dict[Tuple, RotaryEmbedding] = {} _ROPE_DICT: dict[tuple, RotaryEmbedding] = {}
def get_rope( def get_rope(
head_size: int, head_size: int,
rotary_dim: int, rotary_dim: int,
max_position: int, max_position: int,
base: Union[int, float], base: int | float,
is_neox_style: bool = True, is_neox_style: bool = True,
rope_scaling: Optional[Dict[str, Any]] = None, rope_scaling: dict[str, Any] | None = None,
dtype: Optional[torch.dtype] = None, dtype: torch.dtype | None = None,
partial_rotary_factor: float = 1.0, partial_rotary_factor: float = 1.0,
) -> RotaryEmbedding: ) -> RotaryEmbedding:
if dtype is None: if dtype is None:
+1 -2
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # 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 # Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/utils.py
"""Utility methods for model layers.""" """Utility methods for model layers."""
from typing import Tuple
import torch import torch
@@ -10,7 +9,7 @@ def get_token_bin_counts_and_mask(
tokens: torch.Tensor, tokens: torch.Tensor,
vocab_size: int, vocab_size: int,
num_seqs: int, num_seqs: int,
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
# Compute the bin counts for the tokens. # Compute the bin counts for the tokens.
# vocab_size + 1 for padding. # vocab_size + 1 for padding.
bin_counts = torch.zeros((num_seqs, vocab_size + 1), bin_counts = torch.zeros((num_seqs, vocab_size + 1),
+2 -3
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
import math import math
from typing import Optional
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -36,7 +35,7 @@ class PatchEmbed(nn.Module):
prefix: str = ""): prefix: str = ""):
super().__init__() super().__init__()
# Convert patch_size to 2-tuple # Convert patch_size to 2-tuple
if isinstance(patch_size, (list, tuple)): if isinstance(patch_size, list | tuple):
if len(patch_size) == 1: if len(patch_size) == 1:
patch_size = (patch_size[0], patch_size[0]) patch_size = (patch_size[0], patch_size[0])
else: else:
@@ -133,7 +132,7 @@ class ModulateProjection(nn.Module):
hidden_size: int, hidden_size: int,
factor: int = 2, factor: int = 2,
act_layer: str = "silu", act_layer: str = "silu",
dtype: Optional[torch.dtype] = None, dtype: torch.dtype | None = None,
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
+11 -11
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from collections.abc import Sequence
from dataclasses import dataclass from dataclasses import dataclass
from typing import List, Optional, Sequence, Tuple
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
@@ -24,7 +24,7 @@ class UnquantizedEmbeddingMethod(QuantizeMethodBase):
def create_weights(self, layer: torch.nn.Module, def create_weights(self, layer: torch.nn.Module,
input_size_per_partition: int, 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, output_size: int, params_dtype: torch.dtype,
**extra_weight_attrs): **extra_weight_attrs):
"""Create weights for embedding layer.""" """Create weights for embedding layer."""
@@ -39,7 +39,7 @@ class UnquantizedEmbeddingMethod(QuantizeMethodBase):
def apply(self, def apply(self,
layer: torch.nn.Module, layer: torch.nn.Module,
x: torch.Tensor, x: torch.Tensor,
bias: Optional[torch.Tensor] = None) -> torch.Tensor: bias: torch.Tensor | None = None) -> torch.Tensor:
return F.linear(x, layer.weight, bias) return F.linear(x, layer.weight, bias)
def embedding(self, layer: torch.nn.Module, 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, input_: torch.Tensor, org_vocab_start_index: int,
org_vocab_end_index: int, num_org_vocab_padding: int, org_vocab_end_index: int, num_org_vocab_padding: int,
added_vocab_start_index: 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 # torch.compile will fuse all of the pointwise ops below
# into a single kernel, making it very fast # into a single kernel, making it very fast
org_vocab_mask = (input_ >= org_vocab_start_index) & (input_ org_vocab_mask = (input_ >= org_vocab_start_index) & (input_
@@ -197,10 +197,10 @@ class VocabParallelEmbedding(torch.nn.Module):
def __init__(self, def __init__(self,
num_embeddings: int, num_embeddings: int,
embedding_dim: int, embedding_dim: int,
params_dtype: Optional[torch.dtype] = None, params_dtype: torch.dtype | None = None,
org_num_embeddings: Optional[int] = None, org_num_embeddings: int | None = None,
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE, padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = ""): prefix: str = ""):
super().__init__() super().__init__()
@@ -296,7 +296,7 @@ class VocabParallelEmbedding(torch.nn.Module):
org_vocab_start_index, org_vocab_end_index, added_vocab_start_index, org_vocab_start_index, org_vocab_end_index, added_vocab_start_index,
added_vocab_end_index) added_vocab_end_index)
def get_sharded_to_full_mapping(self) -> Optional[List[int]]: def get_sharded_to_full_mapping(self) -> list[int] | None:
"""Get a mapping that can be used to reindex the gathered """Get a mapping that can be used to reindex the gathered
logits for sampling. logits for sampling.
@@ -310,9 +310,9 @@ class VocabParallelEmbedding(torch.nn.Module):
if self.tp_size < 2: if self.tp_size < 2:
return None return None
base_embeddings: List[int] = [] base_embeddings: list[int] = []
added_embeddings: List[int] = [] added_embeddings: list[int] = []
padding: List[int] = [] padding: list[int] = []
for tp_rank in range(self.tp_size): for tp_rank in range(self.tp_size):
shard_indices = self._get_indices(self.num_embeddings_padded, shard_indices = self._get_indices(self.num_embeddings_padded,
self.org_vocab_size_padded, self.org_vocab_size_padded,
+2 -3
View File
@@ -11,7 +11,7 @@ from logging import Logger
from logging.config import dictConfig from logging.config import dictConfig
from os import path from os import path
from types import MethodType from types import MethodType
from typing import Any, Optional, cast from typing import Any, cast
import fastvideo.v1.envs as envs import fastvideo.v1.envs as envs
@@ -278,8 +278,7 @@ def _trace_calls(log_path, root_dir, frame, event, arg=None):
return partial(_trace_calls, log_path, root_dir) return partial(_trace_calls, log_path, root_dir)
def enable_trace_function_call(log_file_path: str, def enable_trace_function_call(log_file_path: str, root_dir: str | None = None):
root_dir: Optional[str] = None):
""" """
Enable tracing of every function call in code under `root_dir`. Enable tracing of every function call in code under `root_dir`.
This is useful for debugging hangs or crashes. This is useful for debugging hangs or crashes.
+7 -7
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any, List, Optional, Tuple, Union from typing import Any
import torch import torch
from torch import nn from torch import nn
@@ -18,7 +18,7 @@ class BaseDiT(nn.Module, ABC):
num_attention_heads: int num_attention_heads: int
num_channels_latents: int num_channels_latents: int
# always supports torch_sdpa # always supports torch_sdpa
_supported_attention_backends: Tuple[ _supported_attention_backends: tuple[
_Backend, ...] = DiTConfig()._supported_attention_backends _Backend, ...] = DiTConfig()._supported_attention_backends
def __init_subclass__(cls) -> None: def __init_subclass__(cls) -> None:
@@ -44,10 +44,10 @@ class BaseDiT(nn.Module, ABC):
@abstractmethod @abstractmethod
def forward(self, def forward(self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]], encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor, timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[ encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
torch.Tensor, List[torch.Tensor]]] = None, | None = None,
guidance=None, guidance=None,
**kwargs) -> torch.Tensor: **kwargs) -> torch.Tensor:
pass pass
@@ -63,7 +63,7 @@ class BaseDiT(nn.Module, ABC):
) )
@property @property
def supported_attention_backends(self) -> Tuple[_Backend, ...]: def supported_attention_backends(self) -> tuple[_Backend, ...]:
return self._supported_attention_backends return self._supported_attention_backends
@@ -81,7 +81,7 @@ class CachableDiT(BaseDiT):
num_attention_heads: int num_attention_heads: int
num_channels_latents: int num_channels_latents: int
# always supports torch_sdpa # always supports torch_sdpa
_supported_attention_backends: Tuple[ _supported_attention_backends: tuple[
_Backend, ...] = DiTConfig()._supported_attention_backends _Backend, ...] = DiTConfig()._supported_attention_backends
def __init__(self, config: DiTConfig, **kwargs) -> None: def __init__(self, config: DiTConfig, **kwargs) -> None:
+9 -11
View File
@@ -1,7 +1,5 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from typing import List, Optional, Tuple, Union
import numpy as np import numpy as np
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -96,8 +94,8 @@ class MMDoubleStreamBlock(nn.Module):
hidden_size: int, hidden_size: int,
num_attention_heads: int, num_attention_heads: int,
mlp_ratio: float, mlp_ratio: float,
dtype: Optional[torch.dtype] = None, dtype: torch.dtype | None = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None, supported_attention_backends: tuple[_Backend, ...] | None = None,
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
@@ -202,7 +200,7 @@ class MMDoubleStreamBlock(nn.Module):
txt: torch.Tensor, txt: torch.Tensor,
vec: torch.Tensor, vec: torch.Tensor,
freqs_cis: tuple, freqs_cis: tuple,
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
# Process modulation vectors # Process modulation vectors
img_mod_outputs = self.img_mod(vec) img_mod_outputs = self.img_mod(vec)
( (
@@ -303,8 +301,8 @@ class MMSingleStreamBlock(nn.Module):
hidden_size: int, hidden_size: int,
num_attention_heads: int, num_attention_heads: int,
mlp_ratio: float = 4.0, mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None, dtype: torch.dtype | None = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None, supported_attention_backends: tuple[_Backend, ...] | None = None,
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
@@ -366,7 +364,7 @@ class MMSingleStreamBlock(nn.Module):
x: torch.Tensor, x: torch.Tensor,
vec: torch.Tensor, vec: torch.Tensor,
txt_len: int, txt_len: int,
freqs_cis: Tuple[torch.Tensor, torch.Tensor], freqs_cis: tuple[torch.Tensor, torch.Tensor],
) -> torch.Tensor: ) -> torch.Tensor:
# Process modulation # Process modulation
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1) mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
@@ -542,10 +540,10 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
# TODO: change output to a dict # TODO: change output to a dict
def forward(self, def forward(self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]], encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor, timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[ encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
torch.Tensor, List[torch.Tensor]]] = None, | None = None,
guidance=None, guidance=None,
**kwargs): **kwargs):
""" """
+17 -16
View File
@@ -10,7 +10,6 @@
# The above copyright notice and this permission notice shall be included in all # The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software. # copies or substantial portions of the Software.
# ============================================================================== # ==============================================================================
from typing import Dict, Optional, Tuple
import torch import torch
from einops import rearrange, repeat from einops import rearrange, repeat
@@ -55,7 +54,7 @@ class PatchEmbed2D(nn.Module):
prefix: str = ""): prefix: str = ""):
super().__init__() super().__init__()
# Convert patch_size to 2-tuple # Convert patch_size to 2-tuple
if isinstance(patch_size, (list, tuple)): if isinstance(patch_size, list | tuple):
if len(patch_size) == 1: if len(patch_size) == 1:
patch_size = (patch_size[0], patch_size[0]) patch_size = (patch_size[0], patch_size[0])
else: else:
@@ -143,7 +142,7 @@ class SelfAttention(nn.Module):
def __init__(self, def __init__(self,
hidden_dim, hidden_dim,
head_dim, head_dim,
rope_split: Tuple[int, int, int] = (64, 32, 32), rope_split: tuple[int, int, int] = (64, 32, 32),
bias: bool = False, bias: bool = False,
with_rope: bool = True, with_rope: bool = True,
with_qk_norm: bool = True, with_qk_norm: bool = True,
@@ -190,8 +189,10 @@ class SelfAttention(nn.Module):
outs = [] outs = []
idx = 0 idx = 0
for (chunk_size, cos_i, sin_i) in zip(self.rope_split, cos_splits, for (chunk_size, cos_i, sin_i) in zip(self.rope_split,
sin_splits): cos_splits,
sin_splits,
strict=False):
# slice the corresponding channels # slice the corresponding channels
x_chunk = x[..., idx:idx + chunk_size] # [B,S,H,chunk_size] x_chunk = x[..., idx:idx + chunk_size] # [B,S,H,chunk_size]
idx += chunk_size idx += chunk_size
@@ -331,8 +332,8 @@ class AdaLayerNormSingle(nn.Module):
def forward( def forward(
self, self,
timestep: torch.Tensor, timestep: torch.Tensor,
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None, added_cond_kwargs: dict[str, torch.Tensor] | None = None,
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
embedded_timestep = self.emb(timestep * self.time_step_rescale) embedded_timestep = self.emb(timestep * self.time_step_rescale)
out, _ = self.linear(self.silu(embedded_timestep)) out, _ = self.linear(self.silu(embedded_timestep))
@@ -377,7 +378,7 @@ class StepVideoTransformerBlock(nn.Module):
dim: int, dim: int,
attention_head_dim: int, attention_head_dim: int,
norm_eps: float = 1e-5, norm_eps: float = 1e-5,
ff_inner_dim: Optional[int] = None, ff_inner_dim: int | None = None,
ff_bias: bool = False, ff_bias: bool = False,
attention_type: str = 'torch'): attention_type: str = 'torch'):
super().__init__() super().__init__()
@@ -417,7 +418,7 @@ class StepVideoTransformerBlock(nn.Module):
kv: torch.Tensor, kv: torch.Tensor,
t_expand: torch.LongTensor, t_expand: torch.LongTensor,
attn_mask=None, attn_mask=None,
rope_positions: Optional[list] = None, rope_positions: list | None = None,
cos_sin=None, cos_sin=None,
mask_strategy=None) -> torch.Tensor: mask_strategy=None) -> torch.Tensor:
@@ -538,7 +539,7 @@ class StepVideoModel(BaseDiT):
return hidden_states return hidden_states
def prepare_attn_mask(self, encoder_attention_mask, encoder_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() kv_seqlens = encoder_attention_mask.sum(dim=1).int()
mask = torch.zeros([len(kv_seqlens), q_seqlen, mask = torch.zeros([len(kv_seqlens), q_seqlen,
max(kv_seqlens)], max(kv_seqlens)],
@@ -593,12 +594,12 @@ class StepVideoModel(BaseDiT):
def forward( def forward(
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None, encoder_hidden_states: torch.Tensor | None = None,
t_expand: Optional[torch.LongTensor] = None, t_expand: torch.LongTensor | None = None,
encoder_hidden_states_2: Optional[torch.Tensor] = None, encoder_hidden_states_2: torch.Tensor | None = None,
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None, added_cond_kwargs: dict[str, torch.Tensor] | None = None,
encoder_attention_mask: Optional[torch.Tensor] = None, encoder_attention_mask: torch.Tensor | None = None,
fps: Optional[torch.Tensor] = None, fps: torch.Tensor | None = None,
return_dict: bool = True, return_dict: bool = True,
mask_strategy=None, mask_strategy=None,
guidance=None, guidance=None,
+10 -11
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
import math import math
from typing import List, Optional, Tuple, Union
import numpy as np import numpy as np
import torch import torch
@@ -53,7 +52,7 @@ class WanTimeTextImageEmbedding(nn.Module):
dim: int, dim: int,
time_freq_dim: int, time_freq_dim: int,
text_embed_dim: int, text_embed_dim: int,
image_embed_dim: Optional[int] = None, image_embed_dim: int | None = None,
): ):
super().__init__() super().__init__()
@@ -76,7 +75,7 @@ class WanTimeTextImageEmbedding(nn.Module):
self, self,
timestep: torch.Tensor, timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: Optional[torch.Tensor] = None, encoder_hidden_states_image: torch.Tensor | None = None,
): ):
temb = self.time_embedder(timestep) temb = self.time_embedder(timestep)
timestep_proj = self.time_modulation(temb) timestep_proj = self.time_modulation(temb)
@@ -173,7 +172,7 @@ class WanI2VCrossAttention(WanSelfAttention):
window_size=(-1, -1), window_size=(-1, -1),
qk_norm=True, qk_norm=True,
eps=1e-6, eps=1e-6,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None supported_attention_backends: tuple[_Backend, ...] | None = None
) -> None: ) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps, super().__init__(dim, num_heads, window_size, qk_norm, eps,
supported_attention_backends) supported_attention_backends)
@@ -222,9 +221,9 @@ class WanTransformerBlock(nn.Module):
qk_norm: str = "rms_norm_across_heads", qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False, cross_attn_norm: bool = False,
eps: float = 1e-6, eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None, added_kv_proj_dim: int | None = None,
supported_attention_backends: Optional[Tuple[_Backend, supported_attention_backends: tuple[_Backend, ...]
...]] = None, | None = None,
prefix: str = ""): prefix: str = ""):
super().__init__() super().__init__()
@@ -292,7 +291,7 @@ class WanTransformerBlock(nn.Module):
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor,
temb: torch.Tensor, temb: torch.Tensor,
freqs_cis: Tuple[torch.Tensor, torch.Tensor], freqs_cis: tuple[torch.Tensor, torch.Tensor],
) -> torch.Tensor: ) -> torch.Tensor:
if hidden_states.dim() == 4: if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1) hidden_states = hidden_states.squeeze(1)
@@ -417,10 +416,10 @@ class WanTransformer3DModel(CachableDiT):
def forward(self, def forward(self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]], encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor, timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[ encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
torch.Tensor, List[torch.Tensor]]] = None, | None = None,
guidance=None, guidance=None,
**kwargs) -> torch.Tensor: **kwargs) -> torch.Tensor:
forward_batch = get_forward_context().forward_batch forward_batch = get_forward_context().forward_batch
+9 -10
View File
@@ -1,5 +1,4 @@
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Optional, Tuple
import torch import torch
from torch import nn from torch import nn
@@ -11,7 +10,7 @@ from fastvideo.v1.platforms import _Backend
class TextEncoder(nn.Module, ABC): class TextEncoder(nn.Module, ABC):
_supported_attention_backends: Tuple[ _supported_attention_backends: tuple[
_Backend, ...] = TextEncoderConfig()._supported_attention_backends _Backend, ...] = TextEncoderConfig()._supported_attention_backends
def __init__(self, config: TextEncoderConfig) -> None: def __init__(self, config: TextEncoderConfig) -> None:
@@ -24,21 +23,21 @@ class TextEncoder(nn.Module, ABC):
@abstractmethod @abstractmethod
def forward(self, def forward(self,
input_ids: Optional[torch.Tensor], input_ids: torch.Tensor | None,
position_ids: Optional[torch.Tensor] = None, position_ids: torch.Tensor | None = None,
attention_mask: Optional[torch.Tensor] = None, attention_mask: torch.Tensor | None = None,
inputs_embeds: Optional[torch.Tensor] = None, inputs_embeds: torch.Tensor | None = None,
output_hidden_states: Optional[bool] = None, output_hidden_states: bool | None = None,
**kwargs) -> BaseEncoderOutput: **kwargs) -> BaseEncoderOutput:
pass pass
@property @property
def supported_attention_backends(self) -> Tuple[_Backend, ...]: def supported_attention_backends(self) -> tuple[_Backend, ...]:
return self._supported_attention_backends return self._supported_attention_backends
class ImageEncoder(nn.Module, ABC): class ImageEncoder(nn.Module, ABC):
_supported_attention_backends: Tuple[ _supported_attention_backends: tuple[
_Backend, ...] = ImageEncoderConfig()._supported_attention_backends _Backend, ...] = ImageEncoderConfig()._supported_attention_backends
def __init__(self, config: ImageEncoderConfig) -> None: def __init__(self, config: ImageEncoderConfig) -> None:
@@ -55,5 +54,5 @@ class ImageEncoder(nn.Module, ABC):
pass pass
@property @property
def supported_attention_backends(self) -> Tuple[_Backend, ...]: def supported_attention_backends(self) -> tuple[_Backend, ...]:
return self._supported_attention_backends 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 # 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 """Minimal implementation of CLIPVisionModel intended to be only used
within a vision language model.""" within a vision language model."""
from typing import Iterable, Optional, Set, Tuple, Union from collections.abc import Iterable
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -91,9 +91,9 @@ class CLIPTextEmbeddings(nn.Module):
def forward( def forward(
self, self,
input_ids: Optional[torch.LongTensor] = None, input_ids: torch.LongTensor | None = None,
position_ids: Optional[torch.LongTensor] = None, position_ids: torch.LongTensor | None = None,
inputs_embeds: Optional[torch.FloatTensor] = None, inputs_embeds: torch.FloatTensor | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
if input_ids is not None: if input_ids is not None:
seq_length = input_ids.shape[-1] seq_length = input_ids.shape[-1]
@@ -128,8 +128,8 @@ class CLIPAttention(nn.Module):
def __init__( def __init__(
self, self,
config: Union[CLIPVisionConfig, CLIPTextConfig], config: CLIPVisionConfig | CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
@@ -209,8 +209,8 @@ class CLIPMLP(nn.Module):
def __init__( def __init__(
self, self,
config: Union[CLIPVisionConfig, CLIPTextConfig], config: CLIPVisionConfig | CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = "", prefix: str = "",
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -239,8 +239,8 @@ class CLIPEncoderLayer(nn.Module):
def __init__( def __init__(
self, self,
config: Union[CLIPTextConfig, CLIPVisionConfig], config: CLIPTextConfig | CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = "", prefix: str = "",
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -284,9 +284,9 @@ class CLIPEncoder(nn.Module):
def __init__( def __init__(
self, self,
config: Union[CLIPVisionConfig, CLIPTextConfig], config: CLIPVisionConfig | CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: Optional[int] = None, num_hidden_layers_override: int | None = None,
prefix: str = "", prefix: str = "",
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -305,8 +305,8 @@ class CLIPEncoder(nn.Module):
]) ])
def forward( def forward(
self, inputs_embeds: torch.Tensor, return_all_hidden_states: bool self, inputs_embeds: torch.Tensor, return_all_hidden_states: bool
) -> Union[torch.Tensor, list[torch.Tensor]]: ) -> torch.Tensor | list[torch.Tensor]:
hidden_states_pool = [inputs_embeds] hidden_states_pool = [inputs_embeds]
hidden_states = inputs_embeds hidden_states = inputs_embeds
@@ -325,8 +325,8 @@ class CLIPTextTransformer(nn.Module):
def __init__(self, def __init__(self,
config: CLIPTextConfig, config: CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: Optional[int] = None, num_hidden_layers_override: int | None = None,
prefix: str = ""): prefix: str = ""):
super().__init__() super().__init__()
self.config = config self.config = config
@@ -348,11 +348,11 @@ class CLIPTextTransformer(nn.Module):
def forward( def forward(
self, self,
input_ids: Optional[torch.Tensor], input_ids: torch.Tensor | None,
position_ids: Optional[torch.Tensor] = None, position_ids: torch.Tensor | None = None,
attention_mask: Optional[torch.Tensor] = None, attention_mask: torch.Tensor | None = None,
inputs_embeds: Optional[torch.Tensor] = None, inputs_embeds: torch.Tensor | None = None,
output_hidden_states: Optional[bool] = None, output_hidden_states: bool | None = None,
) -> BaseEncoderOutput: ) -> BaseEncoderOutput:
r""" r"""
Returns: Returns:
@@ -440,11 +440,11 @@ class CLIPTextModel(TextEncoder):
def forward( def forward(
self, self,
input_ids: Optional[torch.Tensor], input_ids: torch.Tensor | None,
position_ids: Optional[torch.Tensor] = None, position_ids: torch.Tensor | None = None,
attention_mask: Optional[torch.Tensor] = None, attention_mask: torch.Tensor | None = None,
inputs_embeds: Optional[torch.Tensor] = None, inputs_embeds: torch.Tensor | None = None,
output_hidden_states: Optional[bool] = None, output_hidden_states: bool | None = None,
**kwargs, **kwargs,
) -> BaseEncoderOutput: ) -> BaseEncoderOutput:
@@ -456,8 +456,8 @@ class CLIPTextModel(TextEncoder):
) )
return outputs return outputs
def load_weights(self, weights: Iterable[Tuple[str, def load_weights(self, weights: Iterable[tuple[str,
torch.Tensor]]) -> Set[str]: torch.Tensor]]) -> set[str]:
# Define mapping for stacked parameters # Define mapping for stacked parameters
stacked_params_mapping = [ stacked_params_mapping = [
@@ -467,7 +467,7 @@ class CLIPTextModel(TextEncoder):
("qkv_proj", "v_proj", "v"), ("qkv_proj", "v_proj", "v"),
] ]
params_dict = dict(self.named_parameters()) params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set() loaded_params: set[str] = set()
for name, loaded_weight in weights: for name, loaded_weight in weights:
# Handle q_proj, k_proj, v_proj -> qkv_proj mapping # Handle q_proj, k_proj, v_proj -> qkv_proj mapping
for param_name, weight_name, shard_id in stacked_params_mapping: for param_name, weight_name, shard_id in stacked_params_mapping:
@@ -498,9 +498,9 @@ class CLIPVisionTransformer(nn.Module):
def __init__( def __init__(
self, self,
config: CLIPVisionConfig, config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: Optional[int] = None, num_hidden_layers_override: int | None = None,
require_post_norm: Optional[bool] = None, require_post_norm: bool | None = None,
prefix: str = "", prefix: str = "",
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -540,7 +540,7 @@ class CLIPVisionTransformer(nn.Module):
def forward( def forward(
self, self,
pixel_values: torch.Tensor, pixel_values: torch.Tensor,
feature_sample_layers: Optional[list[int]] = None, feature_sample_layers: list[int] | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
hidden_states = self.embeddings(pixel_values) hidden_states = self.embeddings(pixel_values)
@@ -582,7 +582,7 @@ class CLIPVisionModel(ImageEncoder):
def forward( def forward(
self, self,
pixel_values: torch.Tensor, pixel_values: torch.Tensor,
feature_sample_layers: Optional[list[int]] = None, feature_sample_layers: list[int] | None = None,
**kwargs, **kwargs,
) -> BaseEncoderOutput: ) -> BaseEncoderOutput:
last_hidden_state = self.vision_model(pixel_values, 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 # (TODO) Add prefix argument for filtering out weights to be loaded
# ref: https://github.com/vllm-project/vllm/pull/7186#discussion_r1734163986 # ref: https://github.com/vllm-project/vllm/pull/7186#discussion_r1734163986
def load_weights(self, weights: Iterable[Tuple[str, def load_weights(self, weights: Iterable[tuple[str,
torch.Tensor]]) -> Set[str]: torch.Tensor]]) -> set[str]:
stacked_params_mapping = [ stacked_params_mapping = [
# (param_name, shard_name, shard_id) # (param_name, shard_name, shard_id)
("qkv_proj", "q_proj", "q"), ("qkv_proj", "q_proj", "q"),
@@ -604,7 +604,7 @@ class CLIPVisionModel(ImageEncoder):
("qkv_proj", "v_proj", "v"), ("qkv_proj", "v_proj", "v"),
] ]
params_dict = dict(self.named_parameters()) params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set() loaded_params: set[str] = set()
layer_count = len(self.vision_model.encoder.layers) layer_count = len(self.vision_model.encoder.layers)
for name, loaded_weight in weights: for name, loaded_weight in weights:
+18 -17
View File
@@ -23,7 +23,8 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
"""Inference-only LLaMA model compatible with HuggingFace weights.""" """Inference-only LLaMA model compatible with HuggingFace weights."""
from typing import Any, Dict, Iterable, Optional, Set, Tuple from collections.abc import Iterable
from typing import Any
import torch import torch
from torch import nn from torch import nn
@@ -52,7 +53,7 @@ class LlamaMLP(nn.Module):
hidden_size: int, hidden_size: int,
intermediate_size: int, intermediate_size: int,
hidden_act: str, hidden_act: str,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
bias: bool = False, bias: bool = False,
prefix: str = "", prefix: str = "",
) -> None: ) -> None:
@@ -92,9 +93,9 @@ class LlamaAttention(nn.Module):
num_heads: int, num_heads: int,
num_kv_heads: int, num_kv_heads: int,
rope_theta: float = 10000, rope_theta: float = 10000,
rope_scaling: Optional[Dict[str, Any]] = None, rope_scaling: dict[str, Any] | None = None,
max_position_embeddings: int = 8192, max_position_embeddings: int = 8192,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
bias: bool = False, bias: bool = False,
bias_o_proj: bool = False, bias_o_proj: bool = False,
prefix: str = "") -> None: prefix: str = "") -> None:
@@ -201,7 +202,7 @@ class LlamaDecoderLayer(nn.Module):
def __init__( def __init__(
self, self,
config: LlamaConfig, config: LlamaConfig,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = "", prefix: str = "",
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -254,8 +255,8 @@ class LlamaDecoderLayer(nn.Module):
self, self,
positions: torch.Tensor, positions: torch.Tensor,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
residual: Optional[torch.Tensor], residual: torch.Tensor | None,
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
# Self Attention # Self Attention
if residual is None: if residual is None:
residual = hidden_states residual = hidden_states
@@ -318,11 +319,11 @@ class LlamaModel(TextEncoder):
def forward( def forward(
self, self,
input_ids: Optional[torch.Tensor], input_ids: torch.Tensor | None,
position_ids: Optional[torch.Tensor] = None, position_ids: torch.Tensor | None = None,
attention_mask: Optional[torch.Tensor] = None, attention_mask: torch.Tensor | None = None,
inputs_embeds: Optional[torch.Tensor] = None, inputs_embeds: torch.Tensor | None = None,
output_hidden_states: Optional[bool] = None, output_hidden_states: bool | None = None,
**kwargs, **kwargs,
) -> BaseEncoderOutput: ) -> BaseEncoderOutput:
output_hidden_states = (output_hidden_states output_hidden_states = (output_hidden_states
@@ -339,7 +340,7 @@ class LlamaModel(TextEncoder):
0, hidden_states.shape[1], 0, hidden_states.shape[1],
device=hidden_states.device).unsqueeze(0) device=hidden_states.device).unsqueeze(0)
all_hidden_states: Optional[Tuple[Any, ...]] = ( all_hidden_states: tuple[Any, ...] | None = (
) if output_hidden_states else None ) if output_hidden_states else None
for layer in self.layers: for layer in self.layers:
if all_hidden_states is not None: if all_hidden_states is not None:
@@ -367,8 +368,8 @@ class LlamaModel(TextEncoder):
return output return output
def load_weights(self, weights: Iterable[Tuple[str, def load_weights(self, weights: Iterable[tuple[str,
torch.Tensor]]) -> Set[str]: torch.Tensor]]) -> set[str]:
stacked_params_mapping = [ stacked_params_mapping = [
# (param_name, shard_name, shard_id) # (param_name, shard_name, shard_id)
(".qkv_proj", ".q_proj", "q"), (".qkv_proj", ".q_proj", "q"),
@@ -378,7 +379,7 @@ class LlamaModel(TextEncoder):
(".gate_up_proj", ".up_proj", 1), (".gate_up_proj", ".up_proj", 1),
] ]
params_dict = dict(self.named_parameters()) params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set() loaded_params: set[str] = set()
for name, loaded_weight in weights: for name, loaded_weight in weights:
if "rotary_emb.inv_freq" in name: if "rotary_emb.inv_freq" in name:
continue continue
@@ -400,7 +401,7 @@ class LlamaModel(TextEncoder):
# continue # continue
if "scale" in name: if "scale" in name:
# Remapping the name of FP8 kv-scale. # Remapping the name of FP8 kv-scale.
kv_scale_name: Optional[str] = maybe_remap_kv_scale_name( kv_scale_name: str | None = maybe_remap_kv_scale_name(
name, params_dict) name, params_dict)
if kv_scale_name is None: if kv_scale_name is None:
continue continue
+8 -9
View File
@@ -13,7 +13,6 @@
# ============================================================================== # ==============================================================================
import os import os
from functools import wraps from functools import wraps
from typing import List, Optional
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -179,10 +178,10 @@ class StepChatTokenizer:
def vocab_size(self): def vocab_size(self):
return self._tokenizer.vocab_size() 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) 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) return self._tokenizer.decode_ids(token_ids)
@@ -347,9 +346,9 @@ class MultiQueryAttention(nn.Module):
def forward( def forward(
self, self,
x: torch.Tensor, x: torch.Tensor,
mask: Optional[torch.Tensor], mask: torch.Tensor | None,
cu_seqlens: Optional[torch.Tensor], cu_seqlens: torch.Tensor | None,
max_seq_len: Optional[torch.Tensor], max_seq_len: torch.Tensor | None,
): ):
seqlen, bsz, dim = x.shape seqlen, bsz, dim = x.shape
xqkv = self.wqkv(x) xqkv = self.wqkv(x)
@@ -471,9 +470,9 @@ class TransformerBlock(nn.Module):
def forward( def forward(
self, self,
x: torch.Tensor, x: torch.Tensor,
mask: Optional[torch.Tensor], mask: torch.Tensor | None,
cu_seqlens: Optional[torch.Tensor], cu_seqlens: torch.Tensor | None,
max_seq_len: Optional[torch.Tensor], max_seq_len: torch.Tensor | None,
): ):
residual = self.attention.forward(self.attention_norm(x), mask, residual = self.attention.forward(self.attention_norm(x), mask,
cu_seqlens, max_seq_len) cu_seqlens, max_seq_len)
+29 -29
View File
@@ -20,8 +20,8 @@
"""PyTorch T5 & UMT5 model.""" """PyTorch T5 & UMT5 model."""
import math import math
from collections.abc import Iterable
from dataclasses import dataclass from dataclasses import dataclass
from typing import Iterable, Optional, Set, Tuple
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
@@ -64,7 +64,7 @@ class T5DenseActDense(nn.Module):
def __init__(self, def __init__(self,
config: T5Config, config: T5Config,
quant_config: Optional[QuantizationConfig] = None): quant_config: QuantizationConfig | None = None):
super().__init__() super().__init__()
self.wi = MergedColumnParallelLinear(config.d_model, [config.d_ff], self.wi = MergedColumnParallelLinear(config.d_model, [config.d_ff],
bias=False) bias=False)
@@ -85,7 +85,7 @@ class T5DenseGatedActDense(nn.Module):
def __init__(self, def __init__(self,
config: T5Config, config: T5Config,
quant_config: Optional[QuantizationConfig] = None): quant_config: QuantizationConfig | None = None):
super().__init__() super().__init__()
self.wi_0 = MergedColumnParallelLinear(config.d_model, [config.d_ff], self.wi_0 = MergedColumnParallelLinear(config.d_model, [config.d_ff],
bias=False, bias=False,
@@ -113,7 +113,7 @@ class T5LayerFF(nn.Module):
def __init__(self, def __init__(self,
config: T5Config, config: T5Config,
quant_config: Optional[QuantizationConfig] = None): quant_config: QuantizationConfig | None = None):
super().__init__() super().__init__()
if config.is_gated_act: if config.is_gated_act:
self.DenseReluDense = T5DenseGatedActDense( self.DenseReluDense = T5DenseGatedActDense(
@@ -155,7 +155,7 @@ class T5Attention(nn.Module):
config: T5Config, config: T5Config,
attn_type: str, attn_type: str,
has_relative_attention_bias=False, has_relative_attention_bias=False,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = ""): prefix: str = ""):
super().__init__() super().__init__()
self.attn_type = attn_type self.attn_type = attn_type
@@ -294,7 +294,7 @@ class T5Attention(nn.Module):
self, self,
hidden_states: torch.Tensor, # (num_tokens, d_model) hidden_states: torch.Tensor, # (num_tokens, d_model)
attention_mask: torch.Tensor, attention_mask: torch.Tensor,
attn_metadata: Optional[AttentionMetadata] = None, attn_metadata: AttentionMetadata | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
bs, seq_len, _ = hidden_states.shape bs, seq_len, _ = hidden_states.shape
num_seqs = bs num_seqs = bs
@@ -344,7 +344,7 @@ class T5LayerSelfAttention(nn.Module):
self, self,
config, config,
has_relative_attention_bias=False, has_relative_attention_bias=False,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
@@ -361,7 +361,7 @@ class T5LayerSelfAttention(nn.Module):
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
attention_mask: torch.Tensor, attention_mask: torch.Tensor,
attn_metadata: Optional[AttentionMetadata] = None, attn_metadata: AttentionMetadata | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
normed_hidden_states = self.layer_norm.forward_native(hidden_states) normed_hidden_states = self.layer_norm.forward_native(hidden_states)
attention_output = self.SelfAttention( attention_output = self.SelfAttention(
@@ -377,7 +377,7 @@ class T5LayerCrossAttention(nn.Module):
def __init__(self, def __init__(self,
config, config,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = ""): prefix: str = ""):
super().__init__() super().__init__()
self.EncDecAttention = T5Attention(config, self.EncDecAttention = T5Attention(config,
@@ -390,7 +390,7 @@ class T5LayerCrossAttention(nn.Module):
def forward( def forward(
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
attn_metadata: Optional[AttentionMetadata] = None, attn_metadata: AttentionMetadata | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
normed_hidden_states = self.layer_norm.forward_native(hidden_states) normed_hidden_states = self.layer_norm.forward_native(hidden_states)
attention_output = self.EncDecAttention( attention_output = self.EncDecAttention(
@@ -407,7 +407,7 @@ class T5Block(nn.Module):
config: T5Config, config: T5Config,
is_decoder: bool, is_decoder: bool,
has_relative_attention_bias=False, has_relative_attention_bias=False,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = ""): prefix: str = ""):
super().__init__() super().__init__()
self.is_decoder = is_decoder self.is_decoder = is_decoder
@@ -431,7 +431,7 @@ class T5Block(nn.Module):
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
attention_mask: torch.Tensor, attention_mask: torch.Tensor,
attn_metadata: Optional[AttentionMetadata] = None, attn_metadata: AttentionMetadata | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
hidden_states = self.layer[0](hidden_states=hidden_states, hidden_states = self.layer[0](hidden_states=hidden_states,
@@ -455,7 +455,7 @@ class T5Stack(nn.Module):
is_decoder: bool, is_decoder: bool,
n_layers: int, n_layers: int,
embed_tokens=None, embed_tokens=None,
quant_config: Optional[QuantizationConfig] = None, quant_config: QuantizationConfig | None = None,
prefix: str = "", prefix: str = "",
is_umt5: bool = False): is_umt5: bool = False):
super().__init__() super().__init__()
@@ -524,11 +524,11 @@ class T5EncoderModel(TextEncoder):
def forward( def forward(
self, self,
input_ids: Optional[torch.Tensor], input_ids: torch.Tensor | None,
position_ids: Optional[torch.Tensor] = None, position_ids: torch.Tensor | None = None,
attention_mask: Optional[torch.Tensor] = None, attention_mask: torch.Tensor | None = None,
inputs_embeds: Optional[torch.Tensor] = None, inputs_embeds: torch.Tensor | None = None,
output_hidden_states: Optional[bool] = None, output_hidden_states: bool | None = None,
**kwargs, **kwargs,
) -> BaseEncoderOutput: ) -> BaseEncoderOutput:
attn_metadata = AttentionMetadata(None) attn_metadata = AttentionMetadata(None)
@@ -540,8 +540,8 @@ class T5EncoderModel(TextEncoder):
return BaseEncoderOutput(last_hidden_state=hidden_states) return BaseEncoderOutput(last_hidden_state=hidden_states)
def load_weights(self, weights: Iterable[Tuple[str, def load_weights(self, weights: Iterable[tuple[str,
torch.Tensor]]) -> Set[str]: torch.Tensor]]) -> set[str]:
stacked_params_mapping = [ stacked_params_mapping = [
# (param_name, shard_name, shard_id) # (param_name, shard_name, shard_id)
(".qkv_proj", ".q", "q"), (".qkv_proj", ".q", "q"),
@@ -549,7 +549,7 @@ class T5EncoderModel(TextEncoder):
(".qkv_proj", ".v", "v"), (".qkv_proj", ".v", "v"),
] ]
params_dict = dict(self.named_parameters()) params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set() loaded_params: set[str] = set()
for name, loaded_weight in weights: for name, loaded_weight in weights:
loaded = False loaded = False
if "decoder" in name or "lm_head" in name: if "decoder" in name or "lm_head" in name:
@@ -611,11 +611,11 @@ class UMT5EncoderModel(TextEncoder):
def forward( def forward(
self, self,
input_ids: Optional[torch.Tensor], input_ids: torch.Tensor | None,
position_ids: Optional[torch.Tensor] = None, position_ids: torch.Tensor | None = None,
attention_mask: Optional[torch.Tensor] = None, attention_mask: torch.Tensor | None = None,
inputs_embeds: Optional[torch.Tensor] = None, inputs_embeds: torch.Tensor | None = None,
output_hidden_states: Optional[bool] = None, output_hidden_states: bool | None = None,
**kwargs, **kwargs,
) -> BaseEncoderOutput: ) -> BaseEncoderOutput:
attn_metadata = AttentionMetadata(None) attn_metadata = AttentionMetadata(None)
@@ -630,8 +630,8 @@ class UMT5EncoderModel(TextEncoder):
attention_mask=attention_mask, attention_mask=attention_mask,
) )
def load_weights(self, weights: Iterable[Tuple[str, def load_weights(self, weights: Iterable[tuple[str,
torch.Tensor]]) -> Set[str]: torch.Tensor]]) -> set[str]:
stacked_params_mapping = [ stacked_params_mapping = [
# (param_name, shard_name, shard_id) # (param_name, shard_name, shard_id)
(".qkv_proj", ".q", "q"), (".qkv_proj", ".q", "q"),
@@ -639,7 +639,7 @@ class UMT5EncoderModel(TextEncoder):
(".qkv_proj", ".v", "v"), (".qkv_proj", ".v", "v"),
] ]
params_dict = dict(self.named_parameters()) params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set() loaded_params: set[str] = set()
for name, loaded_weight in weights: for name, loaded_weight in weights:
loaded = False loaded = False
if "decoder" in name or "lm_head" in name: 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 # 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 abc import ABC, abstractmethod
from typing import Generic, Optional, TypeVar, Union from typing import Generic, TypeVar
import torch import torch
from transformers import PretrainedConfig from transformers import PretrainedConfig
@@ -48,9 +48,9 @@ class VisionEncoderInfo(ABC, Generic[_C]):
def resolve_visual_encoder_outputs( def resolve_visual_encoder_outputs(
encoder_outputs: Union[torch.Tensor, list[torch.Tensor]], encoder_outputs: torch.Tensor | list[torch.Tensor],
feature_sample_layers: Optional[list[int]], feature_sample_layers: list[int] | None,
post_layer_norm: Optional[torch.nn.LayerNorm], post_layer_norm: torch.nn.LayerNorm | None,
max_possible_layers: int, max_possible_layers: int,
) -> torch.Tensor: ) -> torch.Tensor:
"""Given the outputs a visual encoder module that may correspond to the """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 json
import os import os
from pathlib import Path from pathlib import Path
from typing import Any, Dict, Optional, Type, Union from typing import Any
from huggingface_hub import snapshot_download from huggingface_hub import snapshot_download
from transformers import AutoConfig, PretrainedConfig from transformers import AutoConfig, PretrainedConfig
from transformers.models.auto.modeling_auto import ( from transformers.models.auto.modeling_auto import (
MODEL_FOR_CAUSAL_LM_MAPPING_NAMES) MODEL_FOR_CAUSAL_LM_MAPPING_NAMES)
_CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = { _CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = {
# ChatGLMConfig.model_type: ChatGLMConfig, # ChatGLMConfig.model_type: ChatGLMConfig,
# DbrxConfig.model_type: DbrxConfig, # DbrxConfig.model_type: DbrxConfig,
# ExaoneConfig.model_type: ExaoneConfig, # ExaoneConfig.model_type: ExaoneConfig,
@@ -50,8 +50,8 @@ def download_from_hf(model_path: str):
def get_hf_config( def get_hf_config(
model: str, model: str,
trust_remote_code: bool, trust_remote_code: bool,
revision: Optional[str] = None, revision: str | None = None,
model_override_args: Optional[dict] = None, model_override_args: dict | None = None,
**kwargs, **kwargs,
): ):
is_gguf = check_gguf_file(model) is_gguf = check_gguf_file(model)
@@ -83,8 +83,8 @@ def get_hf_config(
def get_diffusers_config( def get_diffusers_config(
model: str, model: str,
fastvideo_args: Optional[dict] = None, fastvideo_args: dict | None = None,
) -> Dict[str, Any]: ) -> dict[str, Any]:
"""Gets a configuration for the given diffusers model. """Gets a configuration for the given diffusers model.
Args: Args:
@@ -104,7 +104,7 @@ def get_diffusers_config(
try: try:
# Load the config directly from the file # Load the config directly from the file
with open(config_file) as f: 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 # TODO(will): apply any overrides from inference args
return config_dict return config_dict
@@ -139,7 +139,7 @@ def attach_additional_stop_token_ids(tokenizer):
tokenizer.additional_stop_token_ids = None tokenizer.additional_stop_token_ids = None
def check_gguf_file(model: Union[str, os.PathLike]) -> bool: def check_gguf_file(model: str | os.PathLike) -> bool:
"""Check if the file is a GGUF model.""" """Check if the file is a GGUF model."""
model = Path(model) model = Path(model)
if not model.is_file(): if not model.is_file():
@@ -6,7 +6,8 @@ import json
import os import os
import time import time
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any, Generator, Iterable, List, Optional, Tuple, cast from collections.abc import Generator, Iterable
from typing import Any, cast
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -105,7 +106,7 @@ class TextEncoderLoader(ComponentLoader):
fall_back_to_pt: bool = True fall_back_to_pt: bool = True
"""Whether .pt weights can be used.""" """Whether .pt weights can be used."""
allow_patterns_overrides: Optional[list[str]] = None allow_patterns_overrides: list[str] | None = None
"""If defined, weights will load exclusively using these patterns.""" """If defined, weights will load exclusively using these patterns."""
counter_before_loading_weights: float = 0.0 counter_before_loading_weights: float = 0.0
@@ -115,8 +116,8 @@ class TextEncoderLoader(ComponentLoader):
self, self,
model_name_or_path: str, model_name_or_path: str,
fall_back_to_pt: bool, fall_back_to_pt: bool,
allow_patterns_overrides: Optional[list[str]], allow_patterns_overrides: list[str] | None,
) -> Tuple[str, List[str], bool]: ) -> tuple[str, list[str], bool]:
"""Prepare weights for the model. """Prepare weights for the model.
If the model is not local, it will be downloaded.""" If the model is not local, it will be downloaded."""
@@ -138,7 +139,7 @@ class TextEncoderLoader(ComponentLoader):
hf_folder = model_name_or_path hf_folder = model_name_or_path
hf_weights_files: List[str] = [] hf_weights_files: list[str] = []
for pattern in allow_patterns: for pattern in allow_patterns:
hf_weights_files += glob.glob(os.path.join(hf_folder, pattern)) hf_weights_files += glob.glob(os.path.join(hf_folder, pattern))
if len(hf_weights_files) > 0: if len(hf_weights_files) > 0:
@@ -161,7 +162,7 @@ class TextEncoderLoader(ComponentLoader):
def _get_weights_iterator( def _get_weights_iterator(
self, source: "Source" 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.""" """Get an iterator for the model weights based on the load format."""
hf_folder, hf_weights_files, use_safetensors = self._prepare_weights( hf_folder, hf_weights_files, use_safetensors = self._prepare_weights(
source.model_or_path, source.fall_back_to_pt, source.model_or_path, source.fall_back_to_pt,
@@ -181,7 +182,7 @@ class TextEncoderLoader(ComponentLoader):
self, self,
model_config: Any, model_config: Any,
model: nn.Module, model: nn.Module,
) -> Generator[Tuple[str, torch.Tensor], None, None]: ) -> Generator[tuple[str, torch.Tensor], None, None]:
primary_weights = TextEncoderLoader.Source( primary_weights = TextEncoderLoader.Source(
model_config.model, model_config.model,
prefix="", prefix="",
+11 -11
View File
@@ -7,9 +7,9 @@
import contextlib import contextlib
import re import re
from collections import defaultdict from collections import defaultdict
from collections.abc import Callable, Generator, Hashable
from itertools import chain from itertools import chain
from typing import (Any, Callable, DefaultDict, Dict, Generator, Hashable, List, from typing import Any
Optional, Tuple, Type)
import torch import torch
from torch import nn from torch import nn
@@ -52,7 +52,7 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
def get_param_names_mapping( 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. 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 # TODO(PY): add compile option
def load_fsdp_model( def load_fsdp_model(
model_cls: Type[nn.Module], model_cls: type[nn.Module],
init_params: Dict[str, Any], init_params: dict[str, Any],
weight_dir_list: List[str], weight_dir_list: list[str],
device: torch.device, device: torch.device,
cpu_offload: bool = False, cpu_offload: bool = False,
default_dtype: Optional[torch.dtype] = torch.bfloat16, default_dtype: torch.dtype | None = torch.bfloat16,
) -> torch.nn.Module: ) -> torch.nn.Module:
with set_default_dtype(default_dtype), torch.device("meta"): with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params) model = model_cls(**init_params)
@@ -129,7 +129,7 @@ def shard_model(
*, *,
cpu_offload: bool, cpu_offload: bool,
reshard_after_forward: bool = True, reshard_after_forward: bool = True,
dp_mesh: Optional[DeviceMesh] = None, dp_mesh: DeviceMesh | None = None,
) -> None: ) -> None:
""" """
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API. 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 # TODO(PY): device mesh for cfg parallel
def load_fsdp_model_from_full_model_state_dict( def load_fsdp_model_from_full_model_state_dict(
model: torch.nn.Module, 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, device: torch.device,
strict: bool = False, strict: bool = False,
cpu_offload: bool = False, cpu_offload: bool = False,
param_names_mapping: Optional[Callable[[str], tuple[str, Any, Any]]] = None, param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None = None,
) -> _IncompatibleKeys: ) -> _IncompatibleKeys:
""" """
Converting full state dict into a sharded state dict 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() meta_sharded_sd = model.state_dict()
sharded_sd = {} 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: for source_param_name, full_tensor in full_sd_iterator:
assert param_names_mapping is not None assert param_names_mapping is not None
target_param_name, merge_index, num_params_to_merge = param_names_mapping( target_param_name, merge_index, num_params_to_merge = param_names_mapping(
+16 -17
View File
@@ -8,8 +8,8 @@ import os
import tempfile import tempfile
import time import time
from collections import defaultdict from collections import defaultdict
from collections.abc import Generator
from pathlib import Path from pathlib import Path
from typing import Generator, List, Optional, Tuple, Union
import filelock import filelock
import huggingface_hub.constants import huggingface_hub.constants
@@ -50,8 +50,7 @@ class DisabledTqdm(tqdm):
super().__init__(*args, **kwargs, disable=True) super().__init__(*args, **kwargs, disable=True)
def get_lock(model_name_or_path: Union[str, Path], def get_lock(model_name_or_path: str | Path, cache_dir: str | None = None):
cache_dir: Optional[str] = None):
lock_dir = cache_dir or temp_dir lock_dir = cache_dir or temp_dir
model_name_or_path = str(model_name_or_path) model_name_or_path = str(model_name_or_path)
os.makedirs(os.path.dirname(lock_dir), exist_ok=True) os.makedirs(os.path.dirname(lock_dir), exist_ok=True)
@@ -77,10 +76,10 @@ def _shared_pointers(tensors):
def download_weights_from_hf( def download_weights_from_hf(
model_name_or_path: str, model_name_or_path: str,
cache_dir: Optional[str], cache_dir: str | None,
allow_patterns: List[str], allow_patterns: list[str],
revision: Optional[str] = None, revision: str | None = None,
ignore_patterns: Optional[Union[str, List[str]]] = None, ignore_patterns: str | list[str] | None = None,
) -> str: ) -> str:
"""Download model weights from Hugging Face Hub. """Download model weights from Hugging Face Hub.
@@ -136,8 +135,8 @@ def download_weights_from_hf(
def download_safetensors_index_file_from_hf( def download_safetensors_index_file_from_hf(
model_name_or_path: str, model_name_or_path: str,
index_file: str, index_file: str,
cache_dir: Optional[str], cache_dir: str | None,
revision: Optional[str] = None, revision: str | None = None,
) -> None: ) -> None:
"""Download hf safetensors index file from Hugging Face Hub. """Download hf safetensors index file from Hugging Face Hub.
@@ -172,9 +171,9 @@ def download_safetensors_index_file_from_hf(
# Passing both of these to the weight loader functionality breaks. # Passing both of these to the weight loader functionality breaks.
# So, we use the index_file to # So, we use the index_file to
# look up which safetensors files should be used. # 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, hf_folder: str,
index_file: str) -> List[str]: index_file: str) -> list[str]:
# model.safetensors.index.json is a mapping from keys in the # model.safetensors.index.json is a mapping from keys in the
# torch state_dict to safetensors file holding that weight. # torch state_dict to safetensors file holding that weight.
index_file_name = os.path.join(hf_folder, index_file) index_file_name = os.path.join(hf_folder, index_file)
@@ -197,7 +196,7 @@ def filter_duplicate_safetensors_files(hf_weights_files: List[str],
def filter_files_not_needed_for_inference( 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. Exclude files that are not needed for inference.
@@ -225,8 +224,8 @@ _BAR_FORMAT = "{desc}: {percentage:3.0f}% Completed | {n_fmt}/{total_fmt} [{elap
def safetensors_weights_iterator( def safetensors_weights_iterator(
hf_weights_files: List[str] hf_weights_files: list[str]
) -> Generator[Tuple[str, torch.Tensor], None, None]: ) -> Generator[tuple[str, torch.Tensor], None, None]:
"""Iterate over the weights in the model safetensor files.""" """Iterate over the weights in the model safetensor files."""
enable_tqdm = not torch.distributed.is_initialized( enable_tqdm = not torch.distributed.is_initialized(
) or torch.distributed.get_rank() == 0 ) or torch.distributed.get_rank() == 0
@@ -243,8 +242,8 @@ def safetensors_weights_iterator(
def pt_weights_iterator( def pt_weights_iterator(
hf_weights_files: List[str] hf_weights_files: list[str]
) -> Generator[Tuple[str, torch.Tensor], None, None]: ) -> Generator[tuple[str, torch.Tensor], None, None]:
"""Iterate over the weights in the model bin/pt files.""" """Iterate over the weights in the model bin/pt files."""
enable_tqdm = not torch.distributed.is_initialized( enable_tqdm = not torch.distributed.is_initialized(
) or torch.distributed.get_rank() == 0 ) or torch.distributed.get_rank() == 0
@@ -280,7 +279,7 @@ def default_weight_loader(param: torch.Tensor,
raise raise
def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> Optional[str]: def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> str | None:
"""Remap the name of FP8 k/v_scale parameters. """Remap the name of FP8 k/v_scale parameters.
This function handles the remapping of FP8 k/v_scale parameter names. This function handles the remapping of FP8 k/v_scale parameter names.
+12 -13
View File
@@ -1,8 +1,9 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/parameter.py # 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 fractions import Fraction
from typing import Any, Callable, Tuple, Union from typing import Any
import torch import torch
from torch.nn import Parameter from torch.nn import Parameter
@@ -112,9 +113,8 @@ class _ColumnvLLMParameter(BasevLLMParameter):
if shard_offset is None or shard_size is None: if shard_offset is None or shard_size is None:
raise ValueError("shard_offset and shard_size must be provided") raise ValueError("shard_offset and shard_size must be provided")
if isinstance( if isinstance(
self, self, PackedColumnParameter
(PackedColumnParameter, | PackedvLLMParameter) and self.packed_dim == self.output_dim:
PackedvLLMParameter)) and self.packed_dim == self.output_dim:
shard_size, shard_offset = self.adjust_shard_indexes_for_packing( shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
shard_offset=shard_offset, shard_size=shard_size) shard_offset=shard_offset, shard_size=shard_size)
@@ -141,9 +141,8 @@ class _ColumnvLLMParameter(BasevLLMParameter):
assert num_heads is not None assert num_heads is not None
if isinstance( if isinstance(
self, self, PackedColumnParameter
(PackedColumnParameter, | PackedvLLMParameter) and self.output_dim == self.packed_dim:
PackedvLLMParameter)) and self.output_dim == self.packed_dim:
shard_size, shard_offset = self.adjust_shard_indexes_for_packing( shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
shard_offset=shard_offset, shard_size=shard_size) shard_offset=shard_offset, shard_size=shard_size)
@@ -230,7 +229,7 @@ class PerTensorScaleParameter(BasevLLMParameter):
self.qkv_idxs = {"q": 0, "k": 1, "v": 2} self.qkv_idxs = {"q": 0, "k": 1, "v": 2}
super().__init__(**kwargs) super().__init__(**kwargs)
def _shard_id_as_int(self, shard_id: Union[str, int]) -> int: def _shard_id_as_int(self, shard_id: str | int) -> int:
if isinstance(shard_id, int): if isinstance(shard_id, int):
return shard_id return shard_id
@@ -255,7 +254,7 @@ class PerTensorScaleParameter(BasevLLMParameter):
super().load_row_parallel_weight(*args, **kwargs) super().load_row_parallel_weight(*args, **kwargs)
def _load_into_shard_id(self, loaded_weight: torch.Tensor, def _load_into_shard_id(self, loaded_weight: torch.Tensor,
shard_id: Union[str, int], **kwargs): shard_id: str | int, **kwargs):
""" """
Slice the parameter data based on the shard id for Slice the parameter data based on the shard id for
loading. loading.
@@ -282,7 +281,7 @@ class PackedColumnParameter(_ColumnvLLMParameter):
for more details on the packed properties. for more details on the packed properties.
""" """
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int, def __init__(self, packed_factor: int | Fraction, packed_dim: int,
**kwargs): **kwargs):
self._packed_factor = packed_factor self._packed_factor = packed_factor
self._packed_dim = packed_dim self._packed_dim = packed_dim
@@ -297,7 +296,7 @@ class PackedColumnParameter(_ColumnvLLMParameter):
return self._packed_factor return self._packed_factor
def adjust_shard_indexes_for_packing(self, shard_size, 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( return _adjust_shard_indexes_for_packing(
shard_size=shard_size, shard_size=shard_size,
shard_offset=shard_offset, shard_offset=shard_offset,
@@ -315,7 +314,7 @@ class PackedvLLMParameter(ModelWeightParameter):
by accounting for packing and optionally, marlin tile size. by accounting for packing and optionally, marlin tile size.
""" """
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int, def __init__(self, packed_factor: int | Fraction, packed_dim: int,
**kwargs): **kwargs):
self._packed_factor = packed_factor self._packed_factor = packed_factor
self._packed_dim = packed_dim self._packed_dim = packed_dim
@@ -404,7 +403,7 @@ def permute_param_layout_(param: BasevLLMParameter, input_dim: int,
def _adjust_shard_indexes_for_packing(shard_size, shard_offset, 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_size = shard_size // packed_factor
shard_offset = shard_offset // packed_factor shard_offset = shard_offset // packed_factor
return shard_size, shard_offset return shard_size, shard_offset
+23 -23
View File
@@ -8,10 +8,10 @@ import subprocess
import sys import sys
import tempfile import tempfile
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from collections.abc import Callable, Set
from dataclasses import dataclass, field from dataclasses import dataclass, field
from functools import lru_cache from functools import lru_cache
from typing import (AbstractSet, Callable, Dict, List, NoReturn, Optional, from typing import NoReturn, TypeVar, cast
Tuple, Type, TypeVar, Union, cast)
import cloudpickle import cloudpickle
from torch import nn from torch import nn
@@ -80,7 +80,7 @@ class _ModelInfo:
architecture: str architecture: str
@staticmethod @staticmethod
def from_model_cls(model: Type[nn.Module]) -> "_ModelInfo": def from_model_cls(model: type[nn.Module]) -> "_ModelInfo":
return _ModelInfo(architecture=model.__name__, ) return _ModelInfo(architecture=model.__name__, )
@@ -91,7 +91,7 @@ class _BaseRegisteredModel(ABC):
raise NotImplementedError raise NotImplementedError
@abstractmethod @abstractmethod
def load_model_cls(self) -> Type[nn.Module]: def load_model_cls(self) -> type[nn.Module]:
raise NotImplementedError raise NotImplementedError
@@ -102,10 +102,10 @@ class _RegisteredModel(_BaseRegisteredModel):
""" """
interfaces: _ModelInfo interfaces: _ModelInfo
model_cls: Type[nn.Module] model_cls: type[nn.Module]
@staticmethod @staticmethod
def from_model_cls(model_cls: Type[nn.Module]): def from_model_cls(model_cls: type[nn.Module]):
return _RegisteredModel( return _RegisteredModel(
interfaces=_ModelInfo.from_model_cls(model_cls), interfaces=_ModelInfo.from_model_cls(model_cls),
model_cls=model_cls, model_cls=model_cls,
@@ -114,7 +114,7 @@ class _RegisteredModel(_BaseRegisteredModel):
def inspect_model_cls(self) -> _ModelInfo: def inspect_model_cls(self) -> _ModelInfo:
return self.interfaces return self.interfaces
def load_model_cls(self) -> Type[nn.Module]: def load_model_cls(self) -> type[nn.Module]:
return self.model_cls return self.model_cls
@@ -159,16 +159,16 @@ class _LazyRegisteredModel(_BaseRegisteredModel):
return _run_in_subprocess( return _run_in_subprocess(
lambda: _ModelInfo.from_model_cls(self.load_model_cls())) 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) 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) @lru_cache(maxsize=128)
def _try_load_model_cls( def _try_load_model_cls(
model_arch: str, model_arch: str,
model: _BaseRegisteredModel, model: _BaseRegisteredModel,
) -> Optional[Type[nn.Module]]: ) -> type[nn.Module] | None:
from fastvideo.v1.platforms import current_platform from fastvideo.v1.platforms import current_platform
current_platform.verify_model_arch(model_arch) current_platform.verify_model_arch(model_arch)
try: try:
@@ -182,7 +182,7 @@ def _try_load_model_cls(
def _try_inspect_model_cls( def _try_inspect_model_cls(
model_arch: str, model_arch: str,
model: _BaseRegisteredModel, model: _BaseRegisteredModel,
) -> Optional[_ModelInfo]: ) -> _ModelInfo | None:
try: try:
return model.inspect_model_cls() return model.inspect_model_cls()
except Exception: except Exception:
@@ -194,15 +194,15 @@ def _try_inspect_model_cls(
@dataclass @dataclass
class _ModelRegistry: class _ModelRegistry:
# Keyed by model_arch # Keyed by model_arch
models: Dict[str, _BaseRegisteredModel] = field(default_factory=dict) models: dict[str, _BaseRegisteredModel] = field(default_factory=dict)
def get_supported_archs(self) -> AbstractSet[str]: def get_supported_archs(self) -> Set[str]:
return self.models.keys() return self.models.keys()
def register_model( def register_model(
self, self,
model_arch: str, model_arch: str,
model_cls: Union[Type[nn.Module], str], model_cls: type[nn.Module] | str,
) -> None: ) -> None:
""" """
Register an external model to be used in vLLM. Register an external model to be used in vLLM.
@@ -232,7 +232,7 @@ class _ModelRegistry:
self.models[model_arch] = model 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() all_supported_archs = self.get_supported_archs()
if any(arch in all_supported_archs for arch in architectures): if any(arch in all_supported_archs for arch in architectures):
@@ -244,13 +244,13 @@ class _ModelRegistry:
f"Model architectures {architectures} are not supported for now. " f"Model architectures {architectures} are not supported for now. "
f"Supported architectures: {all_supported_archs}") f"Supported architectures: {all_supported_archs}")
def _try_load_model_cls(self, model_arch: str) -> Optional[Type[nn.Module]]: def _try_load_model_cls(self, model_arch: str) -> type[nn.Module] | None:
if model_arch not in self.models: if model_arch not in self.models:
return None return None
return _try_load_model_cls(model_arch, self.models[model_arch]) return _try_load_model_cls(model_arch, self.models[model_arch])
def _try_inspect_model_cls(self, model_arch: str) -> Optional[_ModelInfo]: def _try_inspect_model_cls(self, model_arch: str) -> _ModelInfo | None:
if model_arch not in self.models: if model_arch not in self.models:
return None return None
@@ -258,8 +258,8 @@ class _ModelRegistry:
def _normalize_archs( def _normalize_archs(
self, self,
architectures: Union[str, List[str]], architectures: str | list[str],
) -> List[str]: ) -> list[str]:
if isinstance(architectures, str): if isinstance(architectures, str):
architectures = [architectures] architectures = [architectures]
if not architectures: if not architectures:
@@ -274,8 +274,8 @@ class _ModelRegistry:
def inspect_model_cls( def inspect_model_cls(
self, self,
architectures: Union[str, List[str]], architectures: str | list[str],
) -> Tuple[_ModelInfo, str]: ) -> tuple[_ModelInfo, str]:
architectures = self._normalize_archs(architectures) architectures = self._normalize_archs(architectures)
for arch in architectures: for arch in architectures:
@@ -287,8 +287,8 @@ class _ModelRegistry:
def resolve_model_cls( def resolve_model_cls(
self, self,
architectures: Union[str, List[str]], architectures: str | list[str],
) -> Tuple[Type[nn.Module], str]: ) -> tuple[type[nn.Module], str]:
architectures = self._normalize_archs(architectures) architectures = self._normalize_archs(architectures)
for arch in architectures: for arch in architectures:
+3 -4
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Optional, Tuple, Union
import torch import torch
from diffusers.utils import BaseOutput from diffusers.utils import BaseOutput
@@ -32,15 +31,15 @@ class BaseScheduler(ABC):
@abstractmethod @abstractmethod
def scale_model_input(self, def scale_model_input(self,
sample: torch.Tensor, sample: torch.Tensor,
timestep: Optional[int] = None) -> torch.Tensor: timestep: int | None = None) -> torch.Tensor:
pass pass
@abstractmethod @abstractmethod
def step( def step(
self, self,
model_output: torch.Tensor, model_output: torch.Tensor,
timestep: Union[int, torch.Tensor], timestep: int | torch.Tensor,
sample: torch.Tensor, sample: torch.Tensor,
return_dict: bool = True, return_dict: bool = True,
) -> Union[BaseOutput, Tuple]: ) -> BaseOutput | tuple:
pass pass
@@ -20,7 +20,7 @@
# ============================================================================== # ==============================================================================
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Optional, Tuple, Union from typing import Any
import torch import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config from diffusers.configuration_utils import ConfigMixin, register_to_config
@@ -75,7 +75,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
shift: float = 1.0, shift: float = 1.0,
reverse: bool = True, reverse: bool = True,
solver: str = "euler", solver: str = "euler",
n_tokens: Optional[int] = None, n_tokens: int | None = None,
**kwargs, **kwargs,
): ):
sigmas = torch.linspace(1, 0, num_train_timesteps + 1) sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
@@ -130,7 +130,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def set_timesteps( def set_timesteps(
self, self,
num_inference_steps: int, num_inference_steps: int,
device: Union[str, torch.device] = None, device: str | torch.device = None,
n_tokens: int = 0, n_tokens: int = 0,
): ):
""" """
@@ -193,7 +193,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def scale_model_input(self, def scale_model_input(self,
sample: torch.Tensor, sample: torch.Tensor,
timestep: Optional[int] = None) -> torch.Tensor: timestep: int | None = None) -> torch.Tensor:
return sample return sample
def sd3_time_shift(self, t: torch.Tensor): def sd3_time_shift(self, t: torch.Tensor):
@@ -202,11 +202,11 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def step( def step(
self, self,
model_output: torch.FloatTensor, model_output: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor], timestep: float | torch.FloatTensor,
sample: torch.FloatTensor, sample: torch.FloatTensor,
return_dict: bool = True, return_dict: bool = True,
**kwargs, **kwargs,
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]: ) -> FlowMatchDiscreteSchedulerOutput | tuple:
""" """
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion 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). process from the learned model outputs (most often the predicted noise).
@@ -232,7 +232,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
returned, otherwise a tuple is returned where the first element is the sample tensor. 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(( raise ValueError((
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to" "Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass" " `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
@@ -23,7 +23,6 @@
# ============================================================================== # ==============================================================================
import math import math
from typing import List, Optional, Tuple, Union
import numpy as np import numpy as np
import torch import torch
@@ -203,7 +202,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
beta_start: float = 0.0001, beta_start: float = 0.0001,
beta_end: float = 0.02, beta_end: float = 0.02,
beta_schedule: str = "linear", beta_schedule: str = "linear",
trained_betas: Optional[Union[np.ndarray, List[float]]] = None, trained_betas: np.ndarray | list[float] | None = None,
solver_order: int = 2, solver_order: int = 2,
prediction_type: str = "epsilon", prediction_type: str = "epsilon",
thresholding: bool = False, thresholding: bool = False,
@@ -212,16 +211,16 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
predict_x0: bool = True, predict_x0: bool = True,
solver_type: str = "bh2", solver_type: str = "bh2",
lower_order_final: bool = True, lower_order_final: bool = True,
disable_corrector: Tuple[int, ...] = (), disable_corrector: tuple[int, ...] = (),
solver_p: SchedulerMixin = None, solver_p: SchedulerMixin = None,
use_karras_sigmas: Optional[bool] = False, use_karras_sigmas: bool | None = False,
use_exponential_sigmas: Optional[bool] = False, use_exponential_sigmas: bool | None = False,
use_beta_sigmas: Optional[bool] = False, use_beta_sigmas: bool | None = False,
use_flow_sigmas: Optional[bool] = False, use_flow_sigmas: bool | None = False,
flow_shift: Optional[float] = 1.0, flow_shift: float | None = 1.0,
timestep_spacing: str = "linspace", timestep_spacing: str = "linspace",
steps_offset: int = 0, steps_offset: int = 0,
final_sigmas_type: Optional[str] = "zero", # "zero", "sigma_min" final_sigmas_type: str | None = "zero", # "zero", "sigma_min"
rescale_betas_zero_snr: bool = False, rescale_betas_zero_snr: bool = False,
): ):
if self.config.use_beta_sigmas and not is_scipy_available(): if self.config.use_beta_sigmas and not is_scipy_available():
@@ -283,21 +282,20 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self.predict_x0 = predict_x0 self.predict_x0 = predict_x0
# setable values # setable values
self.num_inference_steps: Optional[int] = None self.num_inference_steps: int | None = None
timesteps = np.linspace(0, timesteps = np.linspace(0,
num_train_timesteps - 1, num_train_timesteps - 1,
num_train_timesteps, num_train_timesteps,
dtype=np.float32)[::-1].copy() dtype=np.float32)[::-1].copy()
self.timesteps = torch.from_numpy(timesteps) self.timesteps = torch.from_numpy(timesteps)
self.model_outputs = [None] * solver_order self.model_outputs = [None] * solver_order
self.timestep_list: List[Union[int, self.timestep_list: list[int | torch.Tensor] = [None] * solver_order
torch.Tensor]] = [None] * solver_order
self.lower_order_nums = 0 self.lower_order_nums = 0
self.disable_corrector = list(disable_corrector) self.disable_corrector = list(disable_corrector)
self.solver_p = solver_p self.solver_p = solver_p
self.last_sample = None self.last_sample = None
self._step_index: Optional[int] = None self._step_index: int | None = None
self._begin_index: Optional[int] = None self._begin_index: int | None = None
self.sigmas = self.sigmas.to( self.sigmas = self.sigmas.to(
"cpu") # to avoid too much CPU/GPU communication "cpu") # to avoid too much CPU/GPU communication
@@ -333,7 +331,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def set_timesteps(self, def set_timesteps(self,
num_inference_steps: int, num_inference_steps: int,
device: Union[str, torch.device] = None): device: str | torch.device = None):
""" """
Sets the discrete timesteps used for the diffusion chain (to be run before inference). Sets the discrete timesteps used for the diffusion chain (to be run before inference).
@@ -537,7 +535,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler._sigma_to_alpha_sigma_t # Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler._sigma_to_alpha_sigma_t
def _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: if self.config.use_flow_sigmas:
alpha_t = 1 - sigma alpha_t = 1 - sigma
sigma_t = sigma sigma_t = sigma
@@ -708,7 +706,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
model_output: torch.Tensor, model_output: torch.Tensor,
*args, *args,
sample: torch.Tensor = None, sample: torch.Tensor = None,
order: Optional[int] = None, order: int | None = None,
**kwargs, **kwargs,
) -> torch.Tensor: ) -> torch.Tensor:
""" """
@@ -808,7 +806,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
R_tensor: torch.Tensor = torch.stack(R) R_tensor: torch.Tensor = torch.stack(R)
b = torch.tensor(b, device=device) b = torch.tensor(b, device=device)
D1s_tensor: Optional[torch.Tensor] = None D1s_tensor: torch.Tensor | None = None
if len(D1s) > 0: if len(D1s) > 0:
D1s_tensor = torch.stack(D1s, dim=1) # (B, K) D1s_tensor = torch.stack(D1s, dim=1) # (B, K)
# for order 2, we use a simplified version # for order 2, we use a simplified version
@@ -842,9 +840,9 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self, self,
this_model_output: torch.Tensor, this_model_output: torch.Tensor,
*args, *args,
last_sample: Optional[torch.Tensor] = None, last_sample: torch.Tensor | None = None,
this_sample: Optional[torch.Tensor] = None, this_sample: torch.Tensor | None = None,
order: Optional[int] = None, order: int | None = None,
**kwargs, **kwargs,
) -> torch.Tensor: ) -> torch.Tensor:
""" """
@@ -950,7 +948,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
R = torch.stack(R) R = torch.stack(R)
b = torch.tensor(b, device=device) b = torch.tensor(b, device=device)
D1s_tensor: Optional[torch.Tensor] = torch.stack( D1s_tensor: torch.Tensor | None = torch.stack(
D1s, dim=1) if len(D1s) > 0 else None D1s, dim=1) if len(D1s) > 0 else None
# for order 1, we use a simplified version # for order 1, we use a simplified version
@@ -1016,10 +1014,10 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def step( def step(
self, self,
model_output: torch.Tensor, model_output: torch.Tensor,
timestep: Union[int, torch.Tensor], timestep: int | torch.Tensor,
sample: torch.Tensor, sample: torch.Tensor,
return_dict: bool = True, return_dict: bool = True,
) -> Union[SchedulerOutput, Tuple]: ) -> SchedulerOutput | tuple:
""" """
Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with
the multistep UniPC. the multistep UniPC.
+5 -5
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/utils.py # Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/utils.py
"""Utils for model executor.""" """Utils for model executor."""
from typing import Any, Dict, List, Optional from typing import Any
import torch import torch
@@ -58,7 +58,7 @@ def set_random_seed(seed: int) -> None:
def set_weight_attrs( def set_weight_attrs(
weight: torch.Tensor, weight: torch.Tensor,
weight_attrs: Optional[Dict[str, Any]], weight_attrs: dict[str, Any] | None,
): ):
"""Set attributes on a weight tensor. """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 - "model.encoder.layers.0.sub.1" -> ValueError
""" """
subnames = layer_name.split(".") subnames = layer_name.split(".")
int_vals: List[int] = [] int_vals: list[int] = []
for subname in subnames: for subname in subnames:
try: try:
int_vals.append(int(subname)) int_vals.append(int(subname))
@@ -121,8 +121,8 @@ def extract_layer_index(layer_name: str) -> int:
def modulate(x: torch.Tensor, def modulate(x: torch.Tensor,
shift: Optional[torch.Tensor] = None, shift: torch.Tensor | None = None,
scale: Optional[torch.Tensor] = None) -> torch.Tensor: scale: torch.Tensor | None = None) -> torch.Tensor:
"""modulate by shift and scale """modulate by shift and scale
Args: Args:
+17 -17
View File
@@ -1,8 +1,9 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from collections.abc import Iterator
from math import prod from math import prod
from typing import Iterator, Optional, Tuple, Union, cast from typing import Optional, cast
import numpy as np import numpy as np
import torch import torch
@@ -51,8 +52,8 @@ class ParallelTiledVAE(ABC):
return cast(int, self.config.spatial_compression_ratio) return cast(int, self.config.spatial_compression_ratio)
@property @property
def scaling_factor(self) -> Union[float, torch.tensor]: def scaling_factor(self) -> float | torch.Tensor:
return cast(Union[float, torch.tensor], self.config.scaling_factor) return cast(float | torch.Tensor, self.config.scaling_factor)
@abstractmethod @abstractmethod
def _encode(self, *args, **kwargs) -> torch.Tensor: def _encode(self, *args, **kwargs) -> torch.Tensor:
@@ -160,7 +161,7 @@ class ParallelTiledVAE(ABC):
def _parallel_data_generator( def _parallel_data_generator(
self, gathered_results, self, gathered_results,
gathered_dim_metadata) -> Iterator[Tuple[torch.Tensor, int]]: gathered_dim_metadata) -> Iterator[tuple[torch.Tensor, int]]:
global_idx = 0 global_idx = 0
for i, per_rank_metadata in enumerate(gathered_dim_metadata): for i, per_rank_metadata in enumerate(gathered_dim_metadata):
_start_shape = 0 _start_shape = 0
@@ -412,16 +413,16 @@ class ParallelTiledVAE(ABC):
def enable_tiling( def enable_tiling(
self, self,
tile_sample_min_height: Optional[int] = None, tile_sample_min_height: int | None = None,
tile_sample_min_width: Optional[int] = None, tile_sample_min_width: int | None = None,
tile_sample_min_num_frames: Optional[int] = None, tile_sample_min_num_frames: int | None = None,
tile_sample_stride_height: Optional[int] = None, tile_sample_stride_height: int | None = None,
tile_sample_stride_width: Optional[int] = None, tile_sample_stride_width: int | None = None,
tile_sample_stride_num_frames: Optional[int] = None, tile_sample_stride_num_frames: int | None = None,
blend_num_frames: Optional[int] = None, blend_num_frames: int | None = None,
use_tiling: Optional[bool] = None, use_tiling: bool | None = None,
use_temporal_tiling: Optional[bool] = None, use_temporal_tiling: bool | None = None,
use_parallel_tiling: Optional[bool] = None, use_parallel_tiling: bool | None = None,
) -> None: ) -> None:
r""" r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
@@ -485,8 +486,7 @@ class DiagonalGaussianDistribution:
device=self.parameters.device, device=self.parameters.device,
dtype=self.parameters.dtype) dtype=self.parameters.dtype)
def sample(self, def sample(self, generator: torch.Generator | None = None) -> torch.Tensor:
generator: Optional[torch.Generator] = None) -> torch.Tensor:
# make sure sample is on the same device as the parameters and has same dtype # make sure sample is on the same device as the parameters and has same dtype
sample = randn_tensor( sample = randn_tensor(
self.mean.shape, self.mean.shape,
@@ -517,7 +517,7 @@ class DiagonalGaussianDistribution:
def nll( def nll(
self, sample: torch.Tensor, self, sample: torch.Tensor,
dims: Tuple[int, ...] = (1, 2, 3)) -> torch.Tensor: dims: tuple[int, ...] = (1, 2, 3)) -> torch.Tensor:
if self.deterministic: if self.deterministic:
return torch.Tensor([0.0]) return torch.Tensor([0.0])
logtwopi = np.log(2.0 * np.pi) logtwopi = np.log(2.0 * np.pi)
+25 -23
View File
@@ -15,8 +15,6 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
from typing import Optional, Tuple, Union
import numpy as np import numpy as np
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -32,7 +30,7 @@ def prepare_causal_attention_mask(
height_width: int, height_width: int,
dtype: torch.dtype, dtype: torch.dtype,
device: torch.device, device: torch.device,
batch_size: Optional[int] = None) -> torch.Tensor: batch_size: int | None = None) -> torch.Tensor:
indices = torch.arange(1, num_frames + 1, dtype=torch.int32, device=device) indices = torch.arange(1, num_frames + 1, dtype=torch.int32, device=device)
indices_blocks = indices.repeat_interleave(height_width) indices_blocks = indices.repeat_interleave(height_width)
x, y = torch.meshgrid(indices_blocks, indices_blocks, indexing="xy") x, y = torch.meshgrid(indices_blocks, indices_blocks, indexing="xy")
@@ -72,7 +70,7 @@ class HunyuanVAEAttention(nn.Module):
def forward(self, def forward(self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor: attention_mask: torch.Tensor | None = None) -> torch.Tensor:
residual = hidden_states residual = hidden_states
batch_size, sequence_length, _ = hidden_states.shape batch_size, sequence_length, _ = hidden_states.shape
@@ -121,10 +119,10 @@ class HunyuanVideoCausalConv3d(nn.Module):
self, self,
in_channels: int, in_channels: int,
out_channels: int, out_channels: int,
kernel_size: Union[int, Tuple[int, int, int]] = 3, kernel_size: int | tuple[int, int, int] = 3,
stride: Union[int, Tuple[int, int, int]] = 1, stride: int | tuple[int, int, int] = 1,
padding: Union[int, Tuple[int, int, int]] = 0, padding: int | tuple[int, int, int] = 0,
dilation: Union[int, Tuple[int, int, int]] = 1, dilation: int | tuple[int, int, int] = 1,
bias: bool = True, bias: bool = True,
pad_mode: str = "replicate", pad_mode: str = "replicate",
) -> None: ) -> None:
@@ -163,11 +161,11 @@ class HunyuanVideoUpsampleCausal3D(nn.Module):
def __init__( def __init__(
self, self,
in_channels: int, in_channels: int,
out_channels: Optional[int] = None, out_channels: int | None = None,
kernel_size: int = 3, kernel_size: int = 3,
stride: int = 1, stride: int = 1,
bias: bool = True, bias: bool = True,
upsample_factor: Tuple[int, ...] = (2, 2, 2), upsample_factor: tuple[int, ...] = (2, 2, 2),
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -213,7 +211,7 @@ class HunyuanVideoDownsampleCausal3D(nn.Module):
def __init__( def __init__(
self, self,
channels: int, channels: int,
out_channels: Optional[int] = None, out_channels: int | None = None,
padding: int = 1, padding: int = 1,
kernel_size: int = 3, kernel_size: int = 3,
bias: bool = True, bias: bool = True,
@@ -239,7 +237,7 @@ class HunyuanVideoResnetBlockCausal3D(nn.Module):
def __init__( def __init__(
self, self,
in_channels: int, in_channels: int,
out_channels: Optional[int] = None, out_channels: int | None = None,
dropout: float = 0.0, dropout: float = 0.0,
groups: int = 32, groups: int = 32,
eps: float = 1e-6, eps: float = 1e-6,
@@ -313,7 +311,7 @@ class HunyuanVideoMidBlock3D(nn.Module):
non_linearity=resnet_act_fn, non_linearity=resnet_act_fn,
) )
] ]
attentions: list[Optional[HunyuanVAEAttention]] = [] attentions: list[HunyuanVAEAttention | None] = []
for _ in range(num_layers): for _ in range(num_layers):
if self.add_attention: if self.add_attention:
@@ -349,7 +347,9 @@ class HunyuanVideoMidBlock3D(nn.Module):
hidden_states = self._gradient_checkpointing_func( hidden_states = self._gradient_checkpointing_func(
self.resnets[0], hidden_states) self.resnets[0], hidden_states)
for attn, resnet in zip(self.attentions, self.resnets[1:]): for attn, resnet in zip(self.attentions,
self.resnets[1:],
strict=False):
if attn is not None: if attn is not None:
batch_size, num_channels, num_frames, height, width = hidden_states.shape batch_size, num_channels, num_frames, height, width = hidden_states.shape
hidden_states = hidden_states.permute(0, 2, 3, 4, hidden_states = hidden_states.permute(0, 2, 3, 4,
@@ -371,7 +371,9 @@ class HunyuanVideoMidBlock3D(nn.Module):
else: else:
hidden_states = self.resnets[0](hidden_states) hidden_states = self.resnets[0](hidden_states)
for attn, resnet in zip(self.attentions, self.resnets[1:]): for attn, resnet in zip(self.attentions,
self.resnets[1:],
strict=False):
if attn is not None: if attn is not None:
batch_size, num_channels, num_frames, height, width = hidden_states.shape batch_size, num_channels, num_frames, height, width = hidden_states.shape
hidden_states = hidden_states.permute(0, 2, 3, 4, hidden_states = hidden_states.permute(0, 2, 3, 4,
@@ -404,7 +406,7 @@ class HunyuanVideoDownBlock3D(nn.Module):
resnet_act_fn: str = "silu", resnet_act_fn: str = "silu",
resnet_groups: int = 32, resnet_groups: int = 32,
add_downsample: bool = True, add_downsample: bool = True,
downsample_stride: Tuple[int, ...] | int = 2, downsample_stride: tuple[int, ...] | int = 2,
downsample_padding: int = 1, downsample_padding: int = 1,
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -466,7 +468,7 @@ class HunyuanVideoUpBlock3D(nn.Module):
resnet_act_fn: str = "silu", resnet_act_fn: str = "silu",
resnet_groups: int = 32, resnet_groups: int = 32,
add_upsample: bool = True, add_upsample: bool = True,
upsample_scale_factor: Tuple[int, ...] = (2, 2, 2), upsample_scale_factor: tuple[int, ...] = (2, 2, 2),
) -> None: ) -> None:
super().__init__() super().__init__()
resnets = [] resnets = []
@@ -525,13 +527,13 @@ class HunyuanVideoEncoder3D(nn.Module):
self, self,
in_channels: int = 3, in_channels: int = 3,
out_channels: int = 3, out_channels: int = 3,
down_block_types: Tuple[str, ...] = ( down_block_types: tuple[str, ...] = (
"HunyuanVideoDownBlock3D", "HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D", "HunyuanVideoDownBlock3D",
"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, layers_per_block: int = 2,
norm_num_groups: int = 32, norm_num_groups: int = 32,
act_fn: str = "silu", act_fn: str = "silu",
@@ -546,7 +548,7 @@ class HunyuanVideoEncoder3D(nn.Module):
block_out_channels[0], block_out_channels[0],
kernel_size=3, kernel_size=3,
stride=1) stride=1)
self.mid_block: Optional[HunyuanVideoMidBlock3D] = None self.mid_block: HunyuanVideoMidBlock3D | None = None
self.down_blocks = nn.ModuleList([]) self.down_blocks = nn.ModuleList([])
output_channel = block_out_channels[0] output_channel = block_out_channels[0]
@@ -649,13 +651,13 @@ class HunyuanVideoDecoder3D(nn.Module):
self, self,
in_channels: int = 3, in_channels: int = 3,
out_channels: int = 3, out_channels: int = 3,
up_block_types: Tuple[str, ...] = ( up_block_types: tuple[str, ...] = (
"HunyuanVideoUpBlock3D", "HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D", "HunyuanVideoUpBlock3D",
"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, layers_per_block: int = 2,
norm_num_groups: int = 32, norm_num_groups: int = 32,
act_fn: str = "silu", act_fn: str = "silu",
@@ -832,7 +834,7 @@ class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
self, self,
sample: torch.Tensor, sample: torch.Tensor,
sample_posterior: bool = False, sample_posterior: bool = False,
generator: Optional[torch.Generator] = None, generator: torch.Generator | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
r""" r"""
Args: Args:
+5 -5
View File
@@ -10,7 +10,7 @@
# The above copyright notice and this permission notice shall be included in all # The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software. # copies or substantial portions of the Software.
# ============================================================================== # ==============================================================================
from typing import Any, List, Optional, Tuple from typing import Any
import torch import torch
from einops import rearrange from einops import rearrange
@@ -102,7 +102,7 @@ def base_conv3d(x,
return out 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 stride_d, stride_h, stride_w = stride
padding_d, padding_h, padding_w = padding padding_d, padding_h, padding_w = padding
dilation_d, dilation_h, dilation_w = 1, 1, 1 dilation_d, dilation_h, dilation_w = 1, 1, 1
@@ -445,8 +445,8 @@ def base_group_norm_with_zero_pad(x,
class CausalConvChannelLast(CausalConv): class CausalConvChannelLast(CausalConv):
time_causal_padding: Tuple[Any, ...] time_causal_padding: tuple[Any, ...]
time_uncausal_padding: Tuple[Any, ...] time_uncausal_padding: tuple[Any, ...]
def __init__(self, chan_in, chan_out, kernel_size, **kwargs) -> None: def __init__(self, chan_in, chan_out, kernel_size, **kwargs) -> None:
super().__init__(chan_in, chan_out, kernel_size, **kwargs) super().__init__(chan_in, chan_out, kernel_size, **kwargs)
@@ -1121,7 +1121,7 @@ class AutoencoderKLStepvideo(nn.Module, ParallelTiledVAE):
self, self,
sample: torch.Tensor, sample: torch.Tensor,
sample_posterior: bool = False, sample_posterior: bool = False,
generator: Optional[torch.Generator] = None, generator: torch.Generator | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
""" """
Args: Args:
+13 -11
View File
@@ -16,7 +16,6 @@
import contextvars import contextvars
from contextlib import contextmanager from contextlib import contextmanager
from typing import Optional, Tuple, Union
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -68,9 +67,9 @@ class WanCausalConv3d(nn.Conv3d):
self, self,
in_channels: int, in_channels: int,
out_channels: int, out_channels: int,
kernel_size: Union[int, Tuple[int, int, int]], kernel_size: int | tuple[int, int, int],
stride: Union[int, Tuple[int, int, int]] = 1, stride: int | tuple[int, int, int] = 1,
padding: Union[int, Tuple[int, int, int]] = 0, padding: int | tuple[int, int, int] = 0,
) -> None: ) -> None:
super().__init__( super().__init__(
in_channels=in_channels, in_channels=in_channels,
@@ -79,9 +78,9 @@ class WanCausalConv3d(nn.Conv3d):
stride=stride, stride=stride,
padding=padding, padding=padding,
) )
self.padding: Tuple[int, int, int] self.padding: tuple[int, int, int]
# Set up causal padding # 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], self.padding[1], self.padding[1],
2 * self.padding[0], 0) 2 * self.padding[0], 0)
self.padding = (0, 0, 0) self.padding = (0, 0, 0)
@@ -434,7 +433,8 @@ class WanMidBlock(nn.Module):
x = self.resnets[0](x) x = self.resnets[0](x)
# Process through attention and residual blocks # Process through attention and residual blocks
for attn, resnet in zip(self.attentions, self.resnets[1:]): for attn, resnet in zip(self.attentions, self.resnets[1:],
strict=False):
if attn is not None: if attn is not None:
x = attn(x) x = attn(x)
@@ -488,7 +488,8 @@ class WanEncoder3d(nn.Module):
# downsample blocks # downsample blocks
self.down_blocks = nn.ModuleList([]) self.down_blocks = nn.ModuleList([])
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): for i, (in_dim,
out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=False)):
# residual (+attention) blocks # residual (+attention) blocks
for _ in range(num_res_blocks): for _ in range(num_res_blocks):
self.down_blocks.append( self.down_blocks.append(
@@ -589,7 +590,7 @@ class WanUpBlock(nn.Module):
out_dim: int, out_dim: int,
num_res_blocks: int, num_res_blocks: int,
dropout: float = 0.0, dropout: float = 0.0,
upsample_mode: Optional[str] = None, upsample_mode: str | None = None,
non_linearity: str = "silu", non_linearity: str = "silu",
): ):
super().__init__() super().__init__()
@@ -687,7 +688,8 @@ class WanDecoder3d(nn.Module):
# upsample blocks # upsample blocks
self.up_blocks = nn.ModuleList([]) self.up_blocks = nn.ModuleList([])
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): for i, (in_dim,
out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=False)):
# residual (+attention) blocks # residual (+attention) blocks
if i > 0: if i > 0:
in_dim = in_dim // 2 in_dim = in_dim // 2
@@ -944,7 +946,7 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
self, self,
sample: torch.Tensor, sample: torch.Tensor,
sample_posterior: bool = False, sample_posterior: bool = False,
generator: Optional[torch.Generator] = None, generator: torch.Generator | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
""" """
Args: Args:
+11 -15
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
import os import os
from typing import Callable, List, Optional, Tuple, Union from collections.abc import Callable
import numpy as np import numpy as np
import PIL.Image import PIL.Image
@@ -29,8 +29,7 @@ else:
} }
def pil_to_numpy( def pil_to_numpy(images: list[PIL.Image.Image] | PIL.Image.Image) -> np.ndarray:
images: Union[List[PIL.Image.Image], PIL.Image.Image]) -> np.ndarray:
r""" r"""
Convert a PIL image or a list of PIL images to NumPy arrays. Convert a PIL image or a list of PIL images to NumPy arrays.
@@ -69,9 +68,7 @@ def numpy_to_pt(images: np.ndarray) -> torch.Tensor:
return images return images
def normalize( def normalize(images: np.ndarray | torch.Tensor) -> np.ndarray | torch.Tensor:
images: Union[np.ndarray,
torch.Tensor]) -> Union[np.ndarray, torch.Tensor]:
r""" r"""
Normalize an image array to [-1,1]. Normalize an image array to [-1,1].
@@ -87,9 +84,8 @@ def normalize(
def load_image( def load_image(
image: Union[str, PIL.Image.Image], image: str | PIL.Image.Image,
convert_method: Optional[Callable[[PIL.Image.Image], convert_method: Callable[[PIL.Image.Image], PIL.Image.Image] | None = None
PIL.Image.Image]] = None
) -> PIL.Image.Image: ) -> PIL.Image.Image:
""" """
Loads `image` to a PIL Image. Loads `image` to a PIL Image.
@@ -132,11 +128,11 @@ def load_image(
def get_default_height_width( def get_default_height_width(
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor], image: PIL.Image.Image | np.ndarray | torch.Tensor,
vae_scale_factor: int, vae_scale_factor: int,
height: Optional[int] = None, height: int | None = None,
width: Optional[int] = None, width: int | None = None,
) -> Tuple[int, int]: ) -> tuple[int, int]:
r""" r"""
Returns the height and width of the image, downscaled to the next integer multiple of `vae_scale_factor`. Returns the height and width of the image, downscaled to the next integer multiple of `vae_scale_factor`.
@@ -179,12 +175,12 @@ def get_default_height_width(
def resize( def resize(
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor], image: PIL.Image.Image | np.ndarray | torch.Tensor,
height: int, height: int,
width: int, width: int,
resize_mode: str = "default", # "default", "fill", "crop" resize_mode: str = "default", # "default", "fill", "crop"
resample: str = "lanczos", resample: str = "lanczos",
) -> Union[PIL.Image.Image, np.ndarray, torch.Tensor]: ) -> PIL.Image.Image | np.ndarray | torch.Tensor:
""" """
Resize image. Resize image.
@@ -8,7 +8,7 @@ This module defines the base class for pipelines that are composed of multiple s
import os import os
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from copy import deepcopy from copy import deepcopy
from typing import Any, Dict, List, Optional, cast from typing import Any, cast
import torch import torch
@@ -33,20 +33,20 @@ class ComposedPipelineBase(ABC):
""" """
is_video_pipeline: bool = False # To be overridden by video pipelines 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 # TODO(will): args should support both inference args and training args
def __init__(self, def __init__(self,
model_path: str, model_path: str,
fastvideo_args: FastVideoArgs, fastvideo_args: FastVideoArgs,
config: Optional[Dict[str, Any]] = None): config: dict[str, Any] | None = None):
""" """
Initialize the pipeline. After __init__, the pipeline should be ready to Initialize the pipeline. After __init__, the pipeline should be ready to
use. The pipeline should be stateless and not hold any batch state. use. The pipeline should be stateless and not hold any batch state.
""" """
self.model_path = model_path self.model_path = model_path
self._stages: List[PipelineStage] = [] self._stages: list[PipelineStage] = []
self._stage_name_mapping: Dict[str, PipelineStage] = {} self._stage_name_mapping: dict[str, PipelineStage] = {}
if self._required_config_modules is None: if self._required_config_modules is None:
raise NotImplementedError( raise NotImplementedError(
@@ -74,16 +74,16 @@ class ComposedPipelineBase(ABC):
def add_module(self, module_name: str, module: Any): def add_module(self, module_name: str, module: Any):
self.modules[module_name] = module 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) model_path = maybe_download_model(self.model_path)
self.model_path = model_path self.model_path = model_path
# fastvideo_args.downloaded_model_path = model_path # fastvideo_args.downloaded_model_path = model_path
logger.info("Model path: %s", model_path) logger.info("Model path: %s", model_path)
config = verify_model_config_and_directory(model_path) config = verify_model_config_and_directory(model_path)
return cast(Dict[str, Any], config) return cast(dict[str, Any], config)
@property @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 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 the diffusers directory and model_index.json file. These modules will be
@@ -101,7 +101,7 @@ class ComposedPipelineBase(ABC):
return self._required_config_modules return self._required_config_modules
@property @property
def stages(self) -> List[PipelineStage]: def stages(self) -> list[PipelineStage]:
""" """
List of stages in the pipeline. List of stages in the pipeline.
""" """
@@ -120,7 +120,7 @@ class ComposedPipelineBase(ABC):
""" """
return 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. Load the modules from the config.
""" """
+35 -35
View File
@@ -8,7 +8,7 @@ in a functional manner, reducing the need for explicit parameter passing.
""" """
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Union from typing import Any
import torch import torch
@@ -30,81 +30,81 @@ class ForwardBatch:
# specific arguments. # specific arguments.
data_type: str data_type: str
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None generator: torch.Generator | list[torch.Generator] | None = None
# Image inputs # Image inputs
image_path: Optional[str] = None image_path: str | None = None
image_embeds: List[torch.Tensor] = field(default_factory=list) image_embeds: list[torch.Tensor] = field(default_factory=list)
# Text inputs # Text inputs
prompt: Optional[Union[str, List[str]]] = None prompt: str | list[str] | None = None
negative_prompt: Optional[Union[str, List[str]]] = None negative_prompt: str | list[str] | None = None
prompt_path: Optional[str] = None prompt_path: str | None = None
output_path: str = "outputs/" output_path: str = "outputs/"
# Primary encoder embeddings # Primary encoder embeddings
prompt_embeds: List[torch.Tensor] = field(default_factory=list) prompt_embeds: list[torch.Tensor] = field(default_factory=list)
negative_prompt_embeds: Optional[List[torch.Tensor]] = None negative_prompt_embeds: list[torch.Tensor] | None = None
prompt_attention_mask: Optional[List[torch.Tensor]] = None prompt_attention_mask: list[torch.Tensor] | None = None
negative_attention_mask: Optional[List[torch.Tensor]] = None negative_attention_mask: list[torch.Tensor] | None = None
clip_embedding_pos: Optional[List[torch.Tensor]] = None clip_embedding_pos: list[torch.Tensor] | None = None
clip_embedding_neg: Optional[List[torch.Tensor]] = None clip_embedding_neg: list[torch.Tensor] | None = None
# Additional text-related parameters # Additional text-related parameters
max_sequence_length: Optional[int] = None max_sequence_length: int | None = None
prompt_template: Optional[Dict[str, Any]] = None prompt_template: dict[str, Any] | None = None
do_classifier_free_guidance: bool = False do_classifier_free_guidance: bool = False
# Batch info # Batch info
batch_size: Optional[int] = None batch_size: int | None = None
num_videos_per_prompt: int = 1 num_videos_per_prompt: int = 1
seed: Optional[int] = None seed: int | None = None
seeds: Optional[List[int]] = None seeds: list[int] | None = None
# Tracking if embeddings are already processed # Tracking if embeddings are already processed
is_prompt_processed: bool = False is_prompt_processed: bool = False
# Latent tensors # Latent tensors
latents: Optional[torch.Tensor] = None latents: torch.Tensor | None = None
noise_pred: Optional[torch.Tensor] = None noise_pred: torch.Tensor | None = None
image_latent: Optional[torch.Tensor] = None image_latent: torch.Tensor | None = None
# Latent dimensions # Latent dimensions
height_latents: Optional[int] = None height_latents: int | None = None
width_latents: Optional[int] = None width_latents: int | None = None
num_frames: int = 1 # Default for image models num_frames: int = 1 # Default for image models
num_frames_round_down: bool = False # Whether to round down num_frames if it's not divisible by num_gpus num_frames_round_down: bool = False # Whether to round down num_frames if it's not divisible by num_gpus
# Original dimensions (before VAE scaling) # Original dimensions (before VAE scaling)
height: Optional[int] = None height: int | None = None
width: Optional[int] = None width: int | None = None
fps: Optional[int] = None fps: int | None = None
# Timesteps # Timesteps
timesteps: Optional[torch.Tensor] = None timesteps: torch.Tensor | None = None
timestep: Optional[Union[torch.Tensor, float, int]] = None timestep: torch.Tensor | float | int | None = None
step_index: Optional[int] = None step_index: int | None = None
# Scheduler parameters # Scheduler parameters
num_inference_steps: int = 50 num_inference_steps: int = 50
guidance_scale: float = 1.0 guidance_scale: float = 1.0
guidance_rescale: float = 0.0 guidance_rescale: float = 0.0
eta: float = 0.0 eta: float = 0.0
sigmas: Optional[List[float]] = None sigmas: list[float] | None = None
n_tokens: Optional[int] = None n_tokens: int | None = None
# Other parameters that may be needed by specific schedulers # Other parameters that may be needed by specific schedulers
extra_step_kwargs: Dict[str, Any] = field(default_factory=dict) extra_step_kwargs: dict[str, Any] = field(default_factory=dict)
# Component modules (populated by the pipeline) # Component modules (populated by the pipeline)
modules: Dict[str, Any] = field(default_factory=dict) modules: dict[str, Any] = field(default_factory=dict)
# Final output (after pipeline completion) # Final output (after pipeline completion)
output: Any = None output: Any = None
# Extra parameters that might be needed by specific pipeline implementations # Extra parameters that might be needed by specific pipeline implementations
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
# Misc # Misc
save_video: bool = True save_video: bool = True
@@ -112,7 +112,7 @@ class ForwardBatch:
# TeaCache parameters # TeaCache parameters
enable_teacache: bool = False enable_teacache: bool = False
teacache_params: Optional[TeaCacheParams | WanTeaCacheParams] = None teacache_params: TeaCacheParams | WanTeaCacheParams | None = None
def __post_init__(self): def __post_init__(self):
"""Initialize dependent fields after dataclass initialization.""" """Initialize dependent fields after dataclass initialization."""
+6 -6
View File
@@ -4,9 +4,9 @@
import importlib import importlib
import pkgutil import pkgutil
from collections.abc import Set
from dataclasses import dataclass, field from dataclasses import dataclass, field
from functools import lru_cache from functools import lru_cache
from typing import AbstractSet, Dict, Optional, Tuple, Type
from fastvideo.v1.logger import init_logger from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
@@ -17,14 +17,14 @@ logger = init_logger(__name__)
@dataclass @dataclass
class _PipelineRegistry: class _PipelineRegistry:
# Keyed by pipeline_arch # Keyed by pipeline_arch
pipelines: Dict[str, Optional[Type[ComposedPipelineBase]]] = field( pipelines: dict[str, type[ComposedPipelineBase]
default_factory=dict) | None] = field(default_factory=dict)
def get_supported_archs(self) -> AbstractSet[str]: def get_supported_archs(self) -> Set[str]:
return self.pipelines.keys() return self.pipelines.keys()
def _try_load_pipeline_cls( def _try_load_pipeline_cls(
self, pipeline_arch: str) -> Optional[Type[ComposedPipelineBase]]: self, pipeline_arch: str) -> type[ComposedPipelineBase] | None:
if pipeline_arch not in self.pipelines: if pipeline_arch not in self.pipelines:
return None return None
@@ -33,7 +33,7 @@ class _PipelineRegistry:
def resolve_pipeline_cls( def resolve_pipeline_cls(
self, self,
architecture: str, architecture: str,
) -> Tuple[Type[ComposedPipelineBase], str]: ) -> tuple[type[ComposedPipelineBase], str]:
if not architecture: if not architecture:
logger.warning("No pipeline architecture is specified") logger.warning("No pipeline architecture is specified")
+6 -5
View File
@@ -5,7 +5,8 @@ Denoising stage for diffusion pipelines.
import importlib.util import importlib.util
import inspect import inspect
from typing import Any, Dict, Iterable, List, Optional from collections.abc import Iterable
from typing import Any
import torch import torch
from einops import rearrange from einops import rearrange
@@ -108,7 +109,7 @@ class DenoisingStage(PipelineStage):
def dict_to_3d_list(mask_strategy, def dict_to_3d_list(mask_strategy,
t_max=50, t_max=50,
l_max=60, l_max=60,
h_max=24) -> List: h_max=24) -> list:
result = [[[None for _ in range(h_max)] for _ in range(l_max)] result = [[[None for _ in range(h_max)] for _ in range(l_max)]
for _ in range(t_max)] for _ in range(t_max)]
if mask_strategy is None: if mask_strategy is None:
@@ -294,7 +295,7 @@ class DenoisingStage(PipelineStage):
return batch return batch
def prepare_extra_func_kwargs(self, func, kwargs) -> Dict[str, Any]: def prepare_extra_func_kwargs(self, func, kwargs) -> dict[str, Any]:
""" """
Prepare extra kwargs for the scheduler step / denoise step. Prepare extra kwargs for the scheduler step / denoise step.
@@ -313,8 +314,8 @@ class DenoisingStage(PipelineStage):
return extra_step_kwargs return extra_step_kwargs
def progress_bar(self, def progress_bar(self,
iterable: Optional[Iterable] = None, iterable: Iterable | None = None,
total: Optional[int] = None) -> tqdm: total: int | None = None) -> tqdm:
""" """
Create a progress bar for the denoising process. Create a progress bar for the denoising process.
+3 -4
View File
@@ -2,7 +2,6 @@
""" """
Encoding stage for diffusion pipelines. Encoding stage for diffusion pipelines.
""" """
from typing import Optional
import PIL.Image import PIL.Image
import torch import torch
@@ -137,7 +136,7 @@ class EncodingStage(PipelineStage):
def retrieve_latents(self, def retrieve_latents(self,
encoder_output: torch.Tensor, encoder_output: torch.Tensor,
generator: Optional[torch.Generator] = None, generator: torch.Generator | None = None,
sample_mode: str = "sample"): sample_mode: str = "sample"):
if sample_mode == "sample": if sample_mode == "sample":
return encoder_output.sample(generator) return encoder_output.sample(generator)
@@ -151,8 +150,8 @@ class EncodingStage(PipelineStage):
self, self,
image: PIL.Image.Image, image: PIL.Image.Image,
vae_scale_factor: int, vae_scale_factor: int,
height: Optional[int] = None, height: int | None = None,
width: Optional[int] = None, width: int | None = None,
resize_mode: str = "default", # "default", "fill", "crop" resize_mode: str = "default", # "default", "fill", "crop"
) -> torch.Tensor: ) -> torch.Tensor:
image = [image] image = [image]
@@ -56,10 +56,12 @@ class TextEncodingStage(PipelineStage):
fastvideo_args.text_encoder_configs) fastvideo_args.text_encoder_configs)
for tokenizer, text_encoder, encoder_config, preprocess_func, postprocess_func in zip( for tokenizer, text_encoder, encoder_config, preprocess_func, postprocess_func in zip(
self.tokenizers, self.text_encoders, self.tokenizers,
self.text_encoders,
fastvideo_args.text_encoder_configs, fastvideo_args.text_encoder_configs,
fastvideo_args.preprocess_text_funcs, fastvideo_args.preprocess_text_funcs,
fastvideo_args.postprocess_text_funcs): fastvideo_args.postprocess_text_funcs,
strict=False):
if fastvideo_args.use_cpu_offload: if fastvideo_args.use_cpu_offload:
text_encoder = text_encoder.to(fastvideo_args.device) text_encoder = text_encoder.to(fastvideo_args.device)
@@ -9,7 +9,7 @@ using the modular pipeline architecture.
import os import os
from copy import deepcopy from copy import deepcopy
from typing import Any, Dict from typing import Any
import torch import torch
from huggingface_hub import hf_hub_download from huggingface_hub import hf_hub_download
@@ -95,7 +95,7 @@ class StepVideoPipeline(ComposedPipelineBase):
)) ))
torch.ops.load_library(lib_path) torch.ops.load_library(lib_path)
def load_modules(self, fastvideo_args: FastVideoArgs) -> Dict[str, Any]: def load_modules(self, fastvideo_args: FastVideoArgs) -> dict[str, Any]:
""" """
Load the modules from the config. Load the modules from the config.
""" """
+2 -2
View File
@@ -3,7 +3,7 @@
import logging import logging
import traceback import traceback
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING
# imported by other files, do not remove # imported by other files, do not remove
from fastvideo.v1.platforms.interface import _Backend # noqa: F401 from fastvideo.v1.platforms.interface import _Backend # noqa: F401
@@ -13,7 +13,7 @@ from fastvideo.v1.utils import resolve_obj_by_qualname
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def cuda_platform_plugin() -> Optional[str]: def cuda_platform_plugin() -> str | None:
is_cuda = False is_cuda = False
try: try:
+11 -10
View File
@@ -5,8 +5,9 @@ pynvml. However, it should not initialize cuda context.
""" """
import os import os
from collections.abc import Callable
from functools import lru_cache, wraps from functools import lru_cache, wraps
from typing import Callable, List, Optional, Tuple, TypeVar, Union from typing import TypeVar
import torch import torch
from typing_extensions import ParamSpec from typing_extensions import ParamSpec
@@ -69,7 +70,7 @@ class CudaPlatformBase(Platform):
@classmethod @classmethod
def get_device_capability(cls, def get_device_capability(cls,
device_id: int = 0) -> Optional[DeviceCapability]: device_id: int = 0) -> DeviceCapability | None:
raise NotImplementedError raise NotImplementedError
@classmethod @classmethod
@@ -81,7 +82,7 @@ class CudaPlatformBase(Platform):
raise NotImplementedError raise NotImplementedError
@classmethod @classmethod
def is_async_output_supported(cls, enforce_eager: Optional[bool]) -> bool: def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
if enforce_eager: if enforce_eager:
logger.warning( logger.warning(
"To see benefits of async output processing, enable CUDA " "To see benefits of async output processing, enable CUDA "
@@ -91,7 +92,7 @@ class CudaPlatformBase(Platform):
return True return True
@classmethod @classmethod
def is_full_nvlink(cls, device_ids: List[int]) -> bool: def is_full_nvlink(cls, device_ids: list[int]) -> bool:
raise NotImplementedError raise NotImplementedError
@classmethod @classmethod
@@ -100,13 +101,13 @@ class CudaPlatformBase(Platform):
@classmethod @classmethod
def get_current_memory_usage(cls, def get_current_memory_usage(cls,
device: Optional[torch.types.Device] = None device: torch.types.Device | None = None
) -> float: ) -> float:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
return float(torch.cuda.max_memory_allocated(device)) return float(torch.cuda.max_memory_allocated(device))
@classmethod @classmethod
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend], def get_attn_backend_cls(cls, selected_backend: _Backend | None,
head_size: int, dtype: torch.dtype) -> str: head_size: int, dtype: torch.dtype) -> str:
# TODO(will): maybe come up with a more general interface for local attention # TODO(will): maybe come up with a more general interface for local attention
# if distributed is False, we always try to use Flash attn # if distributed is False, we always try to use Flash attn
@@ -250,7 +251,7 @@ class NvmlCudaPlatform(CudaPlatformBase):
@lru_cache(maxsize=8) @lru_cache(maxsize=8)
@with_nvml_context @with_nvml_context
def get_device_capability(cls, def get_device_capability(cls,
device_id: int = 0) -> Optional[DeviceCapability]: device_id: int = 0) -> DeviceCapability | None:
try: try:
physical_device_id = device_id_to_physical_device_id(device_id) physical_device_id = device_id_to_physical_device_id(device_id)
handle = pynvml.nvmlDeviceGetHandleByIndex(physical_device_id) handle = pynvml.nvmlDeviceGetHandleByIndex(physical_device_id)
@@ -264,7 +265,7 @@ class NvmlCudaPlatform(CudaPlatformBase):
@with_nvml_context @with_nvml_context
def has_device_capability( def has_device_capability(
cls, cls,
capability: Union[Tuple[int, int], int], capability: tuple[int, int] | int,
device_id: int = 0, device_id: int = 0,
) -> bool: ) -> bool:
try: try:
@@ -297,7 +298,7 @@ class NvmlCudaPlatform(CudaPlatformBase):
@classmethod @classmethod
@with_nvml_context @with_nvml_context
def is_full_nvlink(cls, physical_device_ids: List[int]) -> bool: def is_full_nvlink(cls, physical_device_ids: list[int]) -> bool:
""" """
query if the set of gpus are fully connected by nvlink (1 hop) query if the set of gpus are fully connected by nvlink (1 hop)
""" """
@@ -362,7 +363,7 @@ class NonNvmlCudaPlatform(CudaPlatformBase):
return int(device_props.total_memory) return int(device_props.total_memory)
@classmethod @classmethod
def is_full_nvlink(cls, physical_device_ids: List[int]) -> bool: def is_full_nvlink(cls, physical_device_ids: list[int]) -> bool:
logger.exception("NVLink detection not possible, as context support was" logger.exception("NVLink detection not possible, as context support was"
" not found. Assuming no NVLink available.") " not found. Assuming no NVLink available.")
return False return False
+7 -7
View File
@@ -3,7 +3,7 @@
import enum import enum
import random import random
from typing import NamedTuple, Optional, Tuple, Union from typing import NamedTuple
import numpy as np import numpy as np
import torch import torch
@@ -87,7 +87,7 @@ class Platform:
return self._enum == PlatformEnum.CUDA return self._enum == PlatformEnum.CUDA
@classmethod @classmethod
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend], def get_attn_backend_cls(cls, selected_backend: _Backend | None,
head_size: int, dtype: torch.dtype) -> str: head_size: int, dtype: torch.dtype) -> str:
"""Get the attention backend class of a device.""" """Get the attention backend class of a device."""
return "" return ""
@@ -96,14 +96,14 @@ class Platform:
def get_device_capability( def get_device_capability(
cls, cls,
device_id: int = 0, device_id: int = 0,
) -> Optional[DeviceCapability]: ) -> DeviceCapability | None:
"""Stateless version of :func:`torch.cuda.get_device_capability`.""" """Stateless version of :func:`torch.cuda.get_device_capability`."""
return None return None
@classmethod @classmethod
def has_device_capability( def has_device_capability(
cls, cls,
capability: Union[Tuple[int, int], int], capability: tuple[int, int] | int,
device_id: int = 0, device_id: int = 0,
) -> bool: ) -> bool:
""" """
@@ -139,7 +139,7 @@ class Platform:
raise NotImplementedError raise NotImplementedError
@classmethod @classmethod
def is_async_output_supported(cls, enforce_eager: Optional[bool]) -> bool: def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
""" """
Check if the current platform supports async output. Check if the current platform supports async output.
""" """
@@ -156,7 +156,7 @@ class Platform:
return torch.inference_mode(mode=True) return torch.inference_mode(mode=True)
@classmethod @classmethod
def seed_everything(cls, seed: Optional[int] = None) -> None: def seed_everything(cls, seed: int | None = None) -> None:
""" """
Set the seed of each random module. Set the seed of each random module.
`torch.manual_seed` will set seed on all devices. `torch.manual_seed` will set seed on all devices.
@@ -193,7 +193,7 @@ class Platform:
@classmethod @classmethod
def get_current_memory_usage(cls, def get_current_memory_usage(cls,
device: Optional[torch.types.Device] = None device: torch.types.Device | None = None
) -> float: ) -> float:
""" """
Return the memory usage in bytes. Return the memory usage in bytes.
+17 -17
View File
@@ -13,10 +13,10 @@ import signal
import sys import sys
import tempfile import tempfile
import traceback import traceback
from collections.abc import Callable
from dataclasses import fields, is_dataclass from dataclasses import fields, is_dataclass
from functools import partial, wraps from functools import partial, wraps
from typing import (Any, Callable, Dict, List, Optional, Tuple, Type, TypeVar, from typing import Any, TypeVar, cast
Union, cast)
import cloudpickle import cloudpickle
import filelock import filelock
@@ -208,7 +208,7 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
return namespace # type: ignore[no-any-return] return namespace # type: ignore[no-any-return]
def _pull_args_from_config(self, args: List[str]) -> List[str]: def _pull_args_from_config(self, args: list[str]) -> list[str]:
"""Method to pull arguments specified in the config file """Method to pull arguments specified in the config file
into the command-line args variable. into the command-line args variable.
@@ -273,7 +273,7 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
return args return args
def _load_config_file(self, file_path: str) -> List[str]: def _load_config_file(self, file_path: str) -> list[str]:
"""Loads a yaml file and returns the key value pairs as a """Loads a yaml file and returns the key value pairs as a
flattened list with argparse like pattern flattened list with argparse like pattern
```yaml ```yaml
@@ -298,9 +298,9 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
"Config file must be of a yaml/yml/json type.\ "Config file must be of a yaml/yml/json type.\
%s supplied", extension) %s supplied", extension)
processed_args: List[str] = [] processed_args: list[str] = []
config: Dict[str, Any] = {} config: dict[str, Any] = {}
try: try:
with open(file_path) as config_file: with open(file_path) as config_file:
config = yaml.safe_load(config_file) config = yaml.safe_load(config_file)
@@ -315,7 +315,7 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
if isinstance(action, StoreBoolean) if isinstance(action, StoreBoolean)
] ]
def process_dict(prefix: str, d: Dict[str, Any]): def process_dict(prefix: str, d: dict[str, Any]):
for key, value in d.items(): for key, value in d.items():
full_key = f"{prefix}.{key}" if prefix else key full_key = f"{prefix}.{key}" if prefix else key
@@ -353,7 +353,7 @@ def get_lock(model_name_or_path: str):
return lock return lock
def warn_for_unimplemented_methods(cls: Type[T]) -> Type[T]: def warn_for_unimplemented_methods(cls: type[T]) -> type[T]:
""" """
A replacement for `abc.ABC`. A replacement for `abc.ABC`.
When we use `abc.ABC`, subclasses will fail to instantiate When we use `abc.ABC`, subclasses will fail to instantiate
@@ -449,7 +449,7 @@ def import_pynvml():
def maybe_download_model(model_path: str, def maybe_download_model(model_path: str,
local_dir: Optional[str] = None, local_dir: str | None = None,
download: bool = True) -> str: download: bool = True) -> str:
""" """
Check if the model path is a Hugging Face Hub model ID and download it if needed. Check if the model path is a Hugging Face Hub model ID and download it if needed.
@@ -485,7 +485,7 @@ def maybe_download_model(model_path: str,
) from e ) from e
def verify_model_config_and_directory(model_path: str) -> Dict[str, Any]: def verify_model_config_and_directory(model_path: str) -> dict[str, Any]:
""" """
Verify that the model directory contains a valid diffusers configuration. Verify that the model directory contains a valid diffusers configuration.
@@ -525,10 +525,10 @@ def verify_model_config_and_directory(model_path: str) -> Dict[str, Any]:
raise ValueError("model_index.json does not contain _diffusers_version") raise ValueError("model_index.json does not contain _diffusers_version")
logger.info("Diffusers version: %s", config["_diffusers_version"]) logger.info("Diffusers version: %s", config["_diffusers_version"])
return cast(Dict[str, Any], config) return cast(dict[str, Any], config)
def maybe_download_model_index(model_name_or_path: str) -> Dict[str, Any]: def maybe_download_model_index(model_name_or_path: str) -> dict[str, Any]:
""" """
Download and extract just the model_index.json for a Hugging Face model. Download and extract just the model_index.json for a Hugging Face model.
@@ -556,7 +556,7 @@ def maybe_download_model_index(model_name_or_path: str) -> Dict[str, Any]:
# Load the model_index.json # Load the model_index.json
with open(model_index_path) as f: with open(model_index_path) as f:
config: Dict[str, Any] = json.load(f) config: dict[str, Any] = json.load(f)
# Verify it has the required fields # Verify it has the required fields
if "_class_name" not in config: if "_class_name" not in config:
@@ -582,7 +582,7 @@ def maybe_download_model_index(model_name_or_path: str) -> Dict[str, Any]:
) from e ) from e
def update_environment_variables(envs: Dict[str, str]): def update_environment_variables(envs: dict[str, str]):
for k, v in envs.items(): for k, v in envs.items():
if k in os.environ and os.environ[k] != v: if k in os.environ and os.environ[k] != v:
logger.warning( logger.warning(
@@ -591,7 +591,7 @@ def update_environment_variables(envs: Dict[str, str]):
os.environ[k] = v os.environ[k] = v
def run_method(obj: Any, method: Union[str, bytes, Callable], args: tuple[Any], def run_method(obj: Any, method: str | bytes | Callable, args: tuple[Any],
kwargs: dict[str, Any]) -> Any: kwargs: dict[str, Any]) -> Any:
""" """
Run a method of an object with the given arguments and keyword arguments. Run a method of an object with the given arguments and keyword arguments.
@@ -613,7 +613,7 @@ def run_method(obj: Any, method: Union[str, bytes, Callable], args: tuple[Any],
return func(*args, **kwargs) return func(*args, **kwargs)
def shallow_asdict(obj) -> Dict[str, Any]: def shallow_asdict(obj) -> dict[str, Any]:
if not is_dataclass(obj): if not is_dataclass(obj):
raise TypeError("Expected dataclass instance") raise TypeError("Expected dataclass instance")
return {f.name: getattr(obj, f.name) for f in fields(obj)} return {f.name: getattr(obj, f.name) for f in fields(obj)}
@@ -637,7 +637,7 @@ def get_exception_traceback() -> str:
class TypeBasedDispatcher: class TypeBasedDispatcher:
def __init__(self, mapping: List[Tuple[Type, Callable]]): def __init__(self, mapping: list[tuple[type, Callable]]):
self._mapping = mapping self._mapping = mapping
def __call__(self, obj: Any): def __call__(self, obj: Any):
+7 -7
View File
@@ -1,6 +1,6 @@
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import (Any, Callable, Dict, List, Optional, Tuple, TypeVar, Union, from collections.abc import Callable
cast) from typing import Any, TypeVar, cast
from fastvideo.v1.fastvideo_args import FastVideoArgs from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.pipelines import ForwardBatch from fastvideo.v1.pipelines import ForwardBatch
@@ -37,7 +37,7 @@ class Executor(ABC):
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs, fastvideo_args: FastVideoArgs,
) -> ForwardBatch: ) -> ForwardBatch:
outputs: List[Dict[str, outputs: list[dict[str,
Any]] = self.collective_rpc("execute_forward", Any]] = self.collective_rpc("execute_forward",
kwargs={ kwargs={
"forward_batch": "forward_batch":
@@ -49,10 +49,10 @@ class Executor(ABC):
@abstractmethod @abstractmethod
def collective_rpc(self, def collective_rpc(self,
method: Union[str, Callable[..., _R]], method: str | Callable[..., _R],
timeout: Optional[float] = None, timeout: float | None = None,
args: Tuple = (), args: tuple = (),
kwargs: Optional[Dict[str, Any]] = None) -> List[_R]: kwargs: dict[str, Any] | None = None) -> list[_R]:
""" """
Execute an RPC call on all workers. Execute an RPC call on all workers.
+3 -3
View File
@@ -5,7 +5,7 @@ import multiprocessing as mp
import os import os
import signal import signal
import sys import sys
from typing import Any, Dict, Optional, TextIO, cast from typing import Any, TextIO, cast
import psutil import psutil
import torch import torch
@@ -92,7 +92,7 @@ class Worker:
output_batch = self.pipeline.forward(forward_batch, self.fastvideo_args) output_batch = self.pipeline.forward(forward_batch, self.fastvideo_args)
return cast(ForwardBatch, output_batch) return cast(ForwardBatch, output_batch)
def shutdown(self) -> Dict[str, Any]: def shutdown(self) -> dict[str, Any]:
"""Gracefully shut down the worker process""" """Gracefully shut down the worker process"""
logger.info("Worker %d shutting down...", logger.info("Worker %d shutting down...",
self.rank, self.rank,
@@ -168,7 +168,7 @@ class Worker:
def init_worker_distributed_environment( def init_worker_distributed_environment(
fastvideo_args: FastVideoArgs, fastvideo_args: FastVideoArgs,
rank: int, rank: int,
distributed_init_method: Optional[str] = None, distributed_init_method: str | None = None,
local_rank: int = -1, local_rank: int = -1,
) -> None: ) -> None:
"""Initialize distributed environment and model parallelism.""" """Initialize distributed environment and model parallelism."""
+6 -5
View File
@@ -4,8 +4,9 @@ import multiprocessing as mp
import os import os
import signal import signal
import time import time
from collections.abc import Callable
from multiprocessing.process import BaseProcess from multiprocessing.process import BaseProcess
from typing import Any, Callable, List, Optional, Union, cast from typing import Any, cast
from fastvideo.v1.fastvideo_args import FastVideoArgs from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger from fastvideo.v1.logger import init_logger
@@ -26,7 +27,7 @@ class MultiprocExecutor(Executor):
# is initialized # is initialized
mp.set_start_method("spawn", force=True) mp.set_start_method("spawn", force=True)
self.workers: List[BaseProcess] = [] self.workers: list[BaseProcess] = []
self.worker_pipes = [] self.worker_pipes = []
# Create pipes and start workers # Create pipes and start workers
@@ -63,10 +64,10 @@ class MultiprocExecutor(Executor):
return cast(ForwardBatch, responses[0]["output_batch"]) return cast(ForwardBatch, responses[0]["output_batch"])
def collective_rpc(self, def collective_rpc(self,
method: Union[str, Callable], method: str | Callable,
timeout: Optional[float] = None, timeout: float | None = None,
args: tuple = (), args: tuple = (),
kwargs: Optional[dict] = None) -> list[Any]: kwargs: dict | None = None) -> list[Any]:
kwargs = kwargs or {} kwargs = kwargs or {}
try: try:
+1 -1
View File
@@ -7,7 +7,7 @@ name = "fastvideo"
version = "0.1.0" version = "0.1.0"
description = "FastVideo" description = "FastVideo"
readme = "README.md" readme = "README.md"
requires-python = ">=3.8" requires-python = ">=3.10"
classifiers = [ classifiers = [
"Programming Language :: Python :: 3", "Programming Language :: Python :: 3",
"License :: OSI Approved :: Apache Software License", "License :: OSI Approved :: Apache Software License",