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