Compare commits
34
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c83432bdea | ||
|
|
e90d7595b1 | ||
|
|
83db422e9c | ||
|
|
dc2a4514e8 | ||
|
|
dd10588fb7 | ||
|
|
8140cb269f | ||
|
|
bb6ae368f2 | ||
|
|
e14d384a71 | ||
|
|
6c528468b5 | ||
|
|
ae5ed0c7e0 | ||
|
|
c37535ab0b | ||
|
|
4e056b92cd | ||
|
|
83348c6ed0 | ||
|
|
691f9d1064 | ||
|
|
796eaf809f | ||
|
|
f51e9d486b | ||
|
|
759f243cf3 | ||
|
|
fb44fbaa1c | ||
|
|
e976583cf4 | ||
|
|
d2db0d475b | ||
|
|
d5ac1e9bee | ||
|
|
1976b23121 | ||
|
|
eafeea4a3f | ||
|
|
bf1fc27989 | ||
|
|
8d99ec3c85 | ||
|
|
b631546e18 | ||
|
|
b2def4b57c | ||
|
|
23fd3ed3c7 | ||
|
|
6e8b11c137 | ||
|
|
fc5a4bc236 | ||
|
|
0f4c8d1360 | ||
|
|
5252d50b25 | ||
|
|
ac07e436bb | ||
|
|
42f902cf23 |
+1
-1
@@ -1,3 +1,3 @@
|
||||
[submodule "csrc/sliding_tile_attention/tk"]
|
||||
path = sta_kernel/thunderkitten/tk
|
||||
path = csrc/sliding_tile_attention/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
@@ -8,5 +8,7 @@ pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-is
|
||||
|
||||
pip install -r requirements-lint.txt
|
||||
|
||||
pip install -r requirements.txt
|
||||
|
||||
# install fastvideo
|
||||
pip install -e .
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from .flash_attn import (DistributedAttention, LocalAttention)
|
||||
|
||||
__all__ = ["DistributedAttention", "LocalAttention"]
|
||||
@@ -0,0 +1,138 @@
|
||||
from itertools import accumulate
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from fastvideo.v1.distributed.communication_op import sequence_model_parallel_all_to_all_4D, sequence_model_parallel_all_gather
|
||||
from fastvideo.v1.distributed.parallel_state import get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size
|
||||
from flash_attn import flash_attn_func, flash_attn_varlen_func
|
||||
|
||||
|
||||
class DistributedAttention(nn.Module):
|
||||
"""Distributed attention module that supports sequence parallelism.
|
||||
|
||||
This class implements a minimal attention operation with support for distributed
|
||||
processing across multiple GPUs using sequence parallelism. The implementation assumes
|
||||
batch_size=1 and no padding tokens for simplicity.
|
||||
|
||||
The sequence parallelism strategy follows the Ulysses paper (https://arxiv.org/abs/2309.14509),
|
||||
which proposes redistributing attention heads across sequence dimension to enable efficient
|
||||
parallel processing of long sequences.
|
||||
|
||||
Args:
|
||||
dropout_rate (float, optional): Dropout probability. Defaults to 0.0.
|
||||
causal (bool, optional): Whether to use causal attention. Defaults to False.
|
||||
softmax_scale (float, optional): Custom scaling factor for attention scores.
|
||||
If None, uses 1/sqrt(head_dim). Defaults to None.
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
dropout_rate: float = 0.0,
|
||||
causal: bool = False,
|
||||
softmax_scale: Optional[float] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dropout_rate = dropout_rate
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
replicated_q: Optional[torch.Tensor] = None,
|
||||
replicated_k: Optional[torch.Tensor] = None,
|
||||
replicated_v: Optional[torch.Tensor] = None,
|
||||
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
"""Forward pass for distributed attention.
|
||||
|
||||
Args:
|
||||
q (torch.Tensor): Query tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
k (torch.Tensor): Key tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
v (torch.Tensor): Value tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
|
||||
replicated_k (Optional[torch.Tensor]): Replicated key tensor
|
||||
replicated_v (Optional[torch.Tensor]): Replicated value tensor
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
|
||||
- o (torch.Tensor): Output tensor after attention for the main sequence
|
||||
- replicated_o (Optional[torch.Tensor]): Output tensor for replicated tokens, if provided
|
||||
"""
|
||||
# Check input shapes
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
|
||||
# assert bs = 1
|
||||
assert q.shape[0] == 1, "Batch size must be 1, and there should be no padding tokens"
|
||||
batch_size, seq_len, num_heads, head_dim = q.shape
|
||||
local_rank = get_sequence_model_parallel_rank()
|
||||
world_size = get_sequence_model_parallel_world_size()
|
||||
|
||||
# Stack QKV
|
||||
qkv = torch.cat([q, k, v], dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
|
||||
# Redistribute heads across sequence dimension
|
||||
qkv = sequence_model_parallel_all_to_all_4D(qkv, scatter_dim=2, gather_dim=1)
|
||||
|
||||
# Concatenate with replicated QKV if provided
|
||||
if replicated_q is not None:
|
||||
assert replicated_k is not None and replicated_v is not None
|
||||
replicated_qkv = torch.cat([replicated_q, replicated_k, replicated_v], dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
heads_per_rank = num_heads // world_size
|
||||
replicated_qkv = replicated_qkv[:, :, local_rank * heads_per_rank:(local_rank + 1) * heads_per_rank]
|
||||
qkv = torch.cat([qkv, replicated_qkv], dim=1)
|
||||
|
||||
q, k, v = qkv.chunk(3, dim=0)
|
||||
# Apply flash attention
|
||||
output = flash_attn_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=self.dropout_rate,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal
|
||||
)
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
if replicated_q is not None:
|
||||
replicated_output = output[:, seq_len*world_size:]
|
||||
output = output[:, :seq_len*world_size]
|
||||
# TODO: make this asynchronous
|
||||
replicated_output = sequence_model_parallel_all_gather(replicated_output, dim=2)
|
||||
output = sequence_model_parallel_all_to_all_4D(output, scatter_dim=1, gather_dim=2)
|
||||
return output, replicated_output
|
||||
|
||||
|
||||
class LocalAttention(nn.Module):
|
||||
def __init__(self, dropout_rate: float = 0.0, causal: bool = False, softmax_scale: Optional[float] = None):
|
||||
super().__init__()
|
||||
self.dropout_rate = dropout_rate
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
|
||||
def forward(self, q, k, v):
|
||||
"""
|
||||
Apply local attention between query, key and value tensors.
|
||||
|
||||
Args:
|
||||
q (torch.Tensor): Query tensor of shape [batch_size, seq_len, num_heads, head_dim]
|
||||
k (torch.Tensor): Key tensor of shape [batch_size, seq_len, num_heads, head_dim]
|
||||
v (torch.Tensor): Value tensor of shape [batch_size, seq_len, num_heads, head_dim]
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor after local attention
|
||||
"""
|
||||
# Check input shapes
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
|
||||
|
||||
# Apply flash attention
|
||||
output = flash_attn_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=self.dropout_rate,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal
|
||||
)
|
||||
|
||||
return output
|
||||
@@ -0,0 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from .communication_op import *
|
||||
from .parallel_state import *
|
||||
from .utils import *
|
||||
@@ -0,0 +1,51 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/communication_op.py
|
||||
|
||||
from typing import Any, Dict, Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
|
||||
from .parallel_state import get_tp_group, get_sp_group
|
||||
|
||||
|
||||
def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
|
||||
"""All-reduce the input tensor across model parallel group."""
|
||||
return get_tp_group().all_reduce(input_)
|
||||
|
||||
|
||||
def tensor_model_parallel_all_gather(input_: torch.Tensor,
|
||||
dim: int = -1) -> torch.Tensor:
|
||||
"""All-gather the input tensor across model parallel group."""
|
||||
return get_tp_group().all_gather(input_, dim)
|
||||
|
||||
|
||||
def tensor_model_parallel_gather(input_: torch.Tensor,
|
||||
dst: int = 0,
|
||||
dim: int = -1) -> Optional[torch.Tensor]:
|
||||
"""Gather the input tensor across model parallel group."""
|
||||
return get_tp_group().gather(input_, dst, dim)
|
||||
|
||||
|
||||
def broadcast_tensor_dict(tensor_dict: Optional[Dict[Any, Union[torch.Tensor,
|
||||
Any]]] = None,
|
||||
src: int = 0):
|
||||
if not torch.distributed.is_initialized():
|
||||
return tensor_dict
|
||||
return get_tp_group().broadcast_tensor_dict(tensor_dict, src)
|
||||
|
||||
|
||||
# TODO: remove model, make it sequence_parallel
|
||||
def sequence_model_parallel_all_to_all_4D(input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1) -> torch.Tensor:
|
||||
"""All-to-all communication of 4D tensors (e.g. QKV matrices) across sequence parallel group."""
|
||||
return get_sp_group().all_to_all_4D(input_, scatter_dim, gather_dim)
|
||||
|
||||
|
||||
def sequence_model_parallel_all_gather(input_: torch.Tensor,
|
||||
dim: int = -1) -> torch.Tensor:
|
||||
"""All-gather the input tensor across model parallel group."""
|
||||
return get_sp_group().all_gather(input_, dim)
|
||||
|
||||
|
||||
@@ -0,0 +1,184 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/base_device_communicator.py
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import ProcessGroup
|
||||
from einops import rearrange
|
||||
|
||||
class DeviceCommunicatorBase:
|
||||
"""
|
||||
Base class for device-specific communicator.
|
||||
It can use the `cpu_group` to initialize the communicator.
|
||||
If the device has PyTorch integration (PyTorch can recognize its
|
||||
communication backend), the `device_group` will also be given.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
cpu_group: ProcessGroup,
|
||||
device: Optional[torch.device] = None,
|
||||
device_group: Optional[ProcessGroup] = None,
|
||||
unique_name: str = ""):
|
||||
self.device = device or torch.device("cpu")
|
||||
self.cpu_group = cpu_group
|
||||
self.device_group = device_group
|
||||
self.unique_name = unique_name
|
||||
self.rank = dist.get_rank(cpu_group)
|
||||
self.world_size = dist.get_world_size(cpu_group)
|
||||
self.ranks = dist.get_process_group_ranks(cpu_group)
|
||||
self.global_rank = dist.get_rank()
|
||||
self.global_world_size = dist.get_world_size()
|
||||
self.rank_in_group = dist.get_group_rank(self.cpu_group,
|
||||
self.global_rank)
|
||||
|
||||
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
dist.all_reduce(input_, group=self.device_group)
|
||||
return input_
|
||||
|
||||
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
if dim < 0:
|
||||
# Convert negative dim to positive.
|
||||
dim += input_.dim()
|
||||
input_size = input_.size()
|
||||
# NOTE: we have to use concat-style all-gather here,
|
||||
# stack-style all-gather has compatibility issues with
|
||||
# torch.compile . see https://github.com/pytorch/pytorch/issues/138795
|
||||
output_size = (input_size[0] * self.world_size, ) + input_size[1:]
|
||||
# Allocate output tensor.
|
||||
output_tensor = torch.empty(output_size,
|
||||
dtype=input_.dtype,
|
||||
device=input_.device)
|
||||
# All-gather.
|
||||
dist.all_gather_into_tensor(output_tensor,
|
||||
input_,
|
||||
group=self.device_group)
|
||||
# Reshape
|
||||
output_tensor = output_tensor.reshape((self.world_size, ) + input_size)
|
||||
output_tensor = output_tensor.movedim(0, dim)
|
||||
output_tensor = output_tensor.reshape(input_size[:dim] +
|
||||
(self.world_size *
|
||||
input_size[dim], ) +
|
||||
input_size[dim + 1:])
|
||||
return output_tensor
|
||||
|
||||
def gather(self,
|
||||
input_: torch.Tensor,
|
||||
dst: int = 0,
|
||||
dim: int = -1) -> Optional[torch.Tensor]:
|
||||
"""
|
||||
NOTE: We assume that the input tensor is on the same device across
|
||||
all the ranks.
|
||||
NOTE: `dst` is the local rank of the destination rank.
|
||||
"""
|
||||
world_size = self.world_size
|
||||
assert -input_.dim() <= dim < input_.dim(), (
|
||||
f"Invalid dim ({dim}) for input tensor with shape {input_.size()}")
|
||||
if dim < 0:
|
||||
# Convert negative dim to positive.
|
||||
dim += input_.dim()
|
||||
|
||||
# Allocate output tensor.
|
||||
if self.rank_in_group == dst:
|
||||
gather_list = [torch.empty_like(input_) for _ in range(world_size)]
|
||||
else:
|
||||
gather_list = None
|
||||
# Gather.
|
||||
torch.distributed.gather(input_,
|
||||
gather_list,
|
||||
dst=self.ranks[dst],
|
||||
group=self.device_group)
|
||||
if self.rank_in_group == dst:
|
||||
output_tensor = torch.cat(gather_list, dim=dim)
|
||||
else:
|
||||
output_tensor = None
|
||||
return output_tensor
|
||||
def all_to_all_4D(self,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1) -> torch.Tensor:
|
||||
"""Specialized all-to-all operation for 4D tensors (e.g., for QKV matrices).
|
||||
|
||||
Args:
|
||||
input_ (torch.Tensor): 4D input tensor to be scattered and gathered.
|
||||
scatter_dim (int, optional): Dimension along which to scatter. Defaults to 2.
|
||||
gather_dim (int, optional): Dimension along which to gather. Defaults to 1.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor after all-to-all operation.
|
||||
"""
|
||||
# Bypass the function if we are using only 1 GPU.
|
||||
if self.world_size == 1:
|
||||
return input_
|
||||
|
||||
assert input_.dim() == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
|
||||
|
||||
if scatter_dim == 2 and gather_dim == 1:
|
||||
# input: (bs, seqlen/P, hc, hs) output: (bs, seqlen, hc/P, hs)
|
||||
bs, shard_seqlen, hc, hs = input_.shape
|
||||
seqlen = shard_seqlen * self.world_size
|
||||
shard_hc = hc // self.world_size
|
||||
|
||||
# Reshape and transpose for scattering
|
||||
input_t = (input_.reshape(bs, shard_seqlen, self.world_size, shard_hc, hs).transpose(0, 2).contiguous())
|
||||
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
|
||||
torch.distributed.all_to_all_single(output, input_t, group=self.device_group)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Reshape and transpose back
|
||||
output = output.reshape(seqlen, bs, shard_hc, hs).transpose(0, 1).contiguous().reshape(bs, seqlen, shard_hc, hs)
|
||||
|
||||
return output
|
||||
|
||||
elif scatter_dim == 1 and gather_dim == 2:
|
||||
# input: (bs, seqlen, hc/P, hs) output: (bs, seqlen/P, hc, hs)
|
||||
bs, seqlen, shard_hc, hs = input_.shape
|
||||
hc = shard_hc * self.world_size
|
||||
shard_seqlen = seqlen // self.world_size
|
||||
|
||||
# Reshape and transpose for scattering
|
||||
input_t = (input_.reshape(bs, self.world_size, shard_seqlen, shard_hc,
|
||||
hs).transpose(0,
|
||||
3).transpose(0,
|
||||
1).contiguous().reshape(self.world_size, shard_hc,
|
||||
shard_seqlen, bs, hs))
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
|
||||
torch.distributed.all_to_all_single(output, input_t, group=self.device_group)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Reshape and transpose back
|
||||
output = output.reshape(hc, shard_seqlen, bs, hs).transpose(0, 2).contiguous().reshape(bs, shard_seqlen, hc, hs)
|
||||
|
||||
return output
|
||||
else:
|
||||
raise RuntimeError("scatter_dim must be 1 or 2 and gather_dim must be 1 or 2")
|
||||
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
|
||||
"""Sends a tensor to the destination rank in a non-blocking way"""
|
||||
"""NOTE: `dst` is the local rank of the destination rank."""
|
||||
if dst is None:
|
||||
dst = (self.rank_in_group + 1) % self.world_size
|
||||
torch.distributed.send(tensor, self.ranks[dst], self.device_group)
|
||||
|
||||
def recv(self,
|
||||
size: torch.Size,
|
||||
dtype: torch.dtype,
|
||||
src: Optional[int] = None) -> torch.Tensor:
|
||||
"""Receives a tensor from the source rank."""
|
||||
"""NOTE: `src` is the local rank of the source rank."""
|
||||
if src is None:
|
||||
src = (self.rank_in_group - 1) % self.world_size
|
||||
|
||||
tensor = torch.empty(size, dtype=dtype, device=self.device)
|
||||
torch.distributed.recv(tensor, self.ranks[src], self.device_group)
|
||||
return tensor
|
||||
|
||||
def destroy(self):
|
||||
pass
|
||||
@@ -0,0 +1,109 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/cuda_communicator.py
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
from .base_device_communicator import DeviceCommunicatorBase
|
||||
|
||||
|
||||
class CudaCommunicator(DeviceCommunicatorBase):
|
||||
|
||||
def __init__(self,
|
||||
cpu_group: ProcessGroup,
|
||||
device: Optional[torch.device] = None,
|
||||
device_group: Optional[ProcessGroup] = None,
|
||||
unique_name: str = ""):
|
||||
super().__init__(cpu_group, device, device_group, unique_name)
|
||||
if "pp" in unique_name:
|
||||
# pipeline parallel does not need custom allreduce
|
||||
use_custom_allreduce = False
|
||||
else:
|
||||
# from vllm.distributed.parallel_state import (
|
||||
# _ENABLE_CUSTOM_ALL_REDUCE)
|
||||
# TODO(will): bring in the custom allreduce from vLLM
|
||||
use_custom_allreduce = False
|
||||
use_pynccl = True
|
||||
|
||||
self.use_pynccl = use_pynccl
|
||||
self.use_custom_allreduce = use_custom_allreduce
|
||||
|
||||
# lazy import to avoid documentation build error
|
||||
# from vllm.distributed.device_communicators.custom_all_reduce import (
|
||||
# CustomAllreduce)
|
||||
from fastvideo.v1.distributed.device_communicators.pynccl import (
|
||||
PyNcclCommunicator)
|
||||
|
||||
self.pynccl_comm: Optional[PyNcclCommunicator] = None
|
||||
if use_pynccl and self.world_size > 1:
|
||||
self.pynccl_comm = PyNcclCommunicator(
|
||||
group=self.cpu_group,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
# TODO(will): bring in the custom allreduce from vLLM
|
||||
self.ca_comm: Optional[CustomAllreduce] = None
|
||||
if use_custom_allreduce and self.world_size > 1:
|
||||
# Initialize a custom fast all-reduce implementation.
|
||||
self.ca_comm = CustomAllreduce(
|
||||
group=self.cpu_group,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
def all_reduce(self, input_):
|
||||
# always try custom allreduce first,
|
||||
# and then pynccl.
|
||||
ca_comm = self.ca_comm
|
||||
if ca_comm is not None and not ca_comm.disabled and \
|
||||
ca_comm.should_custom_ar(input_):
|
||||
out = ca_comm.custom_all_reduce(input_)
|
||||
assert out is not None
|
||||
return out
|
||||
pynccl_comm = self.pynccl_comm
|
||||
assert pynccl_comm is not None
|
||||
out = pynccl_comm.all_reduce(input_)
|
||||
if out is None:
|
||||
# fall back to the default all-reduce using PyTorch.
|
||||
# this usually happens during testing.
|
||||
# when we run the model, allreduce only happens for the TP
|
||||
# group, where we always have either custom allreduce or pynccl.
|
||||
out = input_.clone()
|
||||
torch.distributed.all_reduce(out, group=self.device_group)
|
||||
return out
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
|
||||
"""Sends a tensor to the destination rank in a non-blocking way"""
|
||||
"""NOTE: `dst` is the local rank of the destination rank."""
|
||||
if dst is None:
|
||||
dst = (self.rank_in_group + 1) % self.world_size
|
||||
|
||||
pynccl_comm = self.pynccl_comm
|
||||
if pynccl_comm is not None and not pynccl_comm.disabled:
|
||||
pynccl_comm.send(tensor, dst)
|
||||
else:
|
||||
torch.distributed.send(tensor, self.ranks[dst], self.device_group)
|
||||
|
||||
def recv(self,
|
||||
size: torch.Size,
|
||||
dtype: torch.dtype,
|
||||
src: Optional[int] = None) -> torch.Tensor:
|
||||
"""Receives a tensor from the source rank."""
|
||||
"""NOTE: `src` is the local rank of the source rank."""
|
||||
if src is None:
|
||||
src = (self.rank_in_group - 1) % self.world_size
|
||||
|
||||
tensor = torch.empty(size, dtype=dtype, device=self.device)
|
||||
pynccl_comm = self.pynccl_comm
|
||||
if pynccl_comm is not None and not pynccl_comm.disabled:
|
||||
pynccl_comm.recv(tensor, src)
|
||||
else:
|
||||
torch.distributed.recv(tensor, self.ranks[src], self.device_group)
|
||||
return tensor
|
||||
|
||||
def destroy(self):
|
||||
if self.pynccl_comm is not None:
|
||||
self.pynccl_comm = None
|
||||
if self.ca_comm is not None:
|
||||
self.ca_comm = None
|
||||
@@ -0,0 +1,218 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/pynccl.py
|
||||
|
||||
from typing import Optional, Union
|
||||
|
||||
# ===================== import region =====================
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import ProcessGroup, ReduceOp
|
||||
|
||||
from fastvideo.v1.distributed.device_communicators.pynccl_wrapper import (
|
||||
NCCLLibrary, buffer_type, cudaStream_t, ncclComm_t, ncclDataTypeEnum,
|
||||
ncclRedOpTypeEnum, ncclUniqueId)
|
||||
from fastvideo.v1.distributed.utils import StatelessProcessGroup
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import current_stream
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PyNcclCommunicator:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
group: Union[ProcessGroup, StatelessProcessGroup],
|
||||
device: Union[int, str, torch.device],
|
||||
library_path: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
group: the process group to work on. If None, it will use the
|
||||
default process group.
|
||||
device: the device to bind the PyNcclCommunicator to. If None,
|
||||
it will be bind to f"cuda:{local_rank}".
|
||||
library_path: the path to the NCCL library. If None, it will
|
||||
use the default library path.
|
||||
It is the caller's responsibility to make sure each communicator
|
||||
is bind to a unique device.
|
||||
"""
|
||||
if not isinstance(group, StatelessProcessGroup):
|
||||
assert dist.is_initialized()
|
||||
assert dist.get_backend(group) != dist.Backend.NCCL, (
|
||||
"PyNcclCommunicator should be attached to a non-NCCL group.")
|
||||
# note: this rank is the rank in the group
|
||||
self.rank = dist.get_rank(group)
|
||||
self.world_size = dist.get_world_size(group)
|
||||
else:
|
||||
self.rank = group.rank
|
||||
self.world_size = group.world_size
|
||||
|
||||
self.group = group
|
||||
|
||||
# if world_size == 1, no need to create communicator
|
||||
if self.world_size == 1:
|
||||
self.available = False
|
||||
self.disabled = True
|
||||
return
|
||||
try:
|
||||
self.nccl = NCCLLibrary(library_path)
|
||||
except Exception:
|
||||
# disable because of missing NCCL library
|
||||
# e.g. in a non-GPU environment
|
||||
self.available = False
|
||||
self.disabled = True
|
||||
return
|
||||
|
||||
self.available = True
|
||||
self.disabled = False
|
||||
|
||||
logger.info("FastVideo is using nccl==%s", self.nccl.ncclGetVersion())
|
||||
|
||||
if self.rank == 0:
|
||||
# get the unique id from NCCL
|
||||
self.unique_id = self.nccl.ncclGetUniqueId()
|
||||
else:
|
||||
# construct an empty unique id
|
||||
self.unique_id = ncclUniqueId()
|
||||
|
||||
if not isinstance(group, StatelessProcessGroup):
|
||||
tensor = torch.ByteTensor(list(self.unique_id.internal))
|
||||
ranks = dist.get_process_group_ranks(group)
|
||||
# arg `src` in `broadcast` is the global rank
|
||||
dist.broadcast(tensor, src=ranks[0], group=group)
|
||||
byte_list = tensor.tolist()
|
||||
for i, byte in enumerate(byte_list):
|
||||
self.unique_id.internal[i] = byte
|
||||
else:
|
||||
self.unique_id = group.broadcast_obj(self.unique_id, src=0)
|
||||
if isinstance(device, int):
|
||||
device = torch.device(f"cuda:{device}")
|
||||
elif isinstance(device, str):
|
||||
device = torch.device(device)
|
||||
# now `device` is a `torch.device` object
|
||||
assert isinstance(device, torch.device)
|
||||
self.device = device
|
||||
# nccl communicator and stream will use this device
|
||||
# `torch.cuda.device` is a context manager that changes the
|
||||
# current cuda device to the specified one
|
||||
with torch.cuda.device(device):
|
||||
self.comm: ncclComm_t = self.nccl.ncclCommInitRank(
|
||||
self.world_size, self.unique_id, self.rank)
|
||||
|
||||
stream = current_stream()
|
||||
# A small all_reduce for warmup.
|
||||
data = torch.zeros(1, device=device)
|
||||
self.all_reduce(data)
|
||||
stream.synchronize()
|
||||
del data
|
||||
|
||||
def all_reduce(self,
|
||||
in_tensor: torch.Tensor,
|
||||
op: ReduceOp = ReduceOp.SUM,
|
||||
stream=None) -> torch.Tensor:
|
||||
if self.disabled:
|
||||
return None
|
||||
# nccl communicator created on a specific device
|
||||
# will only work on tensors on the same device
|
||||
# otherwise it will cause "illegal memory access"
|
||||
assert in_tensor.device == self.device, (
|
||||
f"this nccl communicator is created to work on {self.device}, "
|
||||
f"but the input tensor is on {in_tensor.device}")
|
||||
|
||||
out_tensor = torch.empty_like(in_tensor)
|
||||
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
self.nccl.ncclAllReduce(buffer_type(in_tensor.data_ptr()),
|
||||
buffer_type(out_tensor.data_ptr()),
|
||||
in_tensor.numel(),
|
||||
ncclDataTypeEnum.from_torch(in_tensor.dtype),
|
||||
ncclRedOpTypeEnum.from_torch(op), self.comm,
|
||||
cudaStream_t(stream.cuda_stream))
|
||||
return out_tensor
|
||||
|
||||
def all_gather(self,
|
||||
output_tensor: torch.Tensor,
|
||||
input_tensor: torch.Tensor,
|
||||
stream=None):
|
||||
if self.disabled:
|
||||
return
|
||||
# nccl communicator created on a specific device
|
||||
# will only work on tensors on the same device
|
||||
# otherwise it will cause "illegal memory access"
|
||||
assert input_tensor.device == self.device, (
|
||||
f"this nccl communicator is created to work on {self.device}, "
|
||||
f"but the input tensor is on {input_tensor.device}")
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
self.nccl.ncclAllGather(
|
||||
buffer_type(input_tensor.data_ptr()),
|
||||
buffer_type(output_tensor.data_ptr()), input_tensor.numel(),
|
||||
ncclDataTypeEnum.from_torch(input_tensor.dtype), self.comm,
|
||||
cudaStream_t(stream.cuda_stream))
|
||||
|
||||
def reduce_scatter(self,
|
||||
output_tensor: torch.Tensor,
|
||||
input_tensor: torch.Tensor,
|
||||
op: ReduceOp = ReduceOp.SUM,
|
||||
stream=None):
|
||||
if self.disabled:
|
||||
return
|
||||
# nccl communicator created on a specific device
|
||||
# will only work on tensors on the same device
|
||||
# otherwise it will cause "illegal memory access"
|
||||
assert input_tensor.device == self.device, (
|
||||
f"this nccl communicator is created to work on {self.device}, "
|
||||
f"but the input tensor is on {input_tensor.device}")
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
self.nccl.ncclReduceScatter(
|
||||
buffer_type(input_tensor.data_ptr()),
|
||||
buffer_type(output_tensor.data_ptr()), output_tensor.numel(),
|
||||
ncclDataTypeEnum.from_torch(input_tensor.dtype),
|
||||
ncclRedOpTypeEnum.from_torch(op), self.comm,
|
||||
cudaStream_t(stream.cuda_stream))
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: int, stream=None):
|
||||
if self.disabled:
|
||||
return
|
||||
assert tensor.device == self.device, (
|
||||
f"this nccl communicator is created to work on {self.device}, "
|
||||
f"but the input tensor is on {tensor.device}")
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
self.nccl.ncclSend(buffer_type(tensor.data_ptr()), tensor.numel(),
|
||||
ncclDataTypeEnum.from_torch(tensor.dtype), dst,
|
||||
self.comm, cudaStream_t(stream.cuda_stream))
|
||||
|
||||
def recv(self, tensor: torch.Tensor, src: int, stream=None):
|
||||
if self.disabled:
|
||||
return
|
||||
assert tensor.device == self.device, (
|
||||
f"this nccl communicator is created to work on {self.device}, "
|
||||
f"but the input tensor is on {tensor.device}")
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
self.nccl.ncclRecv(buffer_type(tensor.data_ptr()), tensor.numel(),
|
||||
ncclDataTypeEnum.from_torch(tensor.dtype), src,
|
||||
self.comm, cudaStream_t(stream.cuda_stream))
|
||||
|
||||
def broadcast(self, tensor: torch.Tensor, src: int, stream=None):
|
||||
if self.disabled:
|
||||
return
|
||||
assert tensor.device == self.device, (
|
||||
f"this nccl communicator is created to work on {self.device}, "
|
||||
f"but the input tensor is on {tensor.device}")
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
if src == self.rank:
|
||||
sendbuff = buffer_type(tensor.data_ptr())
|
||||
# NCCL requires the sender also to have a receive buffer
|
||||
recvbuff = buffer_type(tensor.data_ptr())
|
||||
else:
|
||||
sendbuff = buffer_type()
|
||||
recvbuff = buffer_type(tensor.data_ptr())
|
||||
self.nccl.ncclBroadcast(sendbuff, recvbuff, tensor.numel(),
|
||||
ncclDataTypeEnum.from_torch(tensor.dtype), src,
|
||||
self.comm, cudaStream_t(stream.cuda_stream))
|
||||
@@ -0,0 +1,342 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/pynccl_wrapper.py
|
||||
|
||||
# This file is a pure Python wrapper for the NCCL library.
|
||||
# The main purpose is to use NCCL combined with CUDA graph.
|
||||
# Before writing this script, we tried the following approach:
|
||||
# 1. We tried to use `cupy`, it calls NCCL correctly, but `cupy` itself
|
||||
# often gets stuck when initializing the NCCL communicator.
|
||||
# 2. We tried to use `torch.distributed`, but `torch.distributed.all_reduce`
|
||||
# contains many other potential cuda APIs, that are not allowed during
|
||||
# capturing the CUDA graph. For further details, please check
|
||||
# https://discuss.pytorch.org/t/pytorch-cudagraph-with-nccl-operation-failed/ .
|
||||
#
|
||||
# Another rejected idea is to write a C/C++ binding for NCCL. It is usually
|
||||
# doable, but we often encounter issues related with nccl versions, and need
|
||||
# to switch between different versions of NCCL. See
|
||||
# https://github.com/NVIDIA/nccl/issues/1234 for more details.
|
||||
# A C/C++ binding is not flexible enough to handle this. It requires
|
||||
# recompilation of the code every time we want to switch between different
|
||||
# versions. This current implementation, with a **pure** Python wrapper, is
|
||||
# more flexible. We can easily switch between different versions of NCCL by
|
||||
# changing the environment variable `VLLM_NCCL_SO_PATH`, or the `so_file`
|
||||
# variable in the code.
|
||||
|
||||
import ctypes
|
||||
import platform
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
from torch.distributed import ReduceOp
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import find_nccl_library
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# === export types and functions from nccl to Python ===
|
||||
# for the original nccl definition, please check
|
||||
# https://github.com/NVIDIA/nccl/blob/master/src/nccl.h.in
|
||||
|
||||
ncclResult_t = ctypes.c_int
|
||||
ncclComm_t = ctypes.c_void_p
|
||||
|
||||
|
||||
class ncclUniqueId(ctypes.Structure):
|
||||
_fields_ = [("internal", ctypes.c_byte * 128)]
|
||||
|
||||
|
||||
cudaStream_t = ctypes.c_void_p
|
||||
buffer_type = ctypes.c_void_p
|
||||
|
||||
ncclDataType_t = ctypes.c_int
|
||||
|
||||
|
||||
class ncclDataTypeEnum:
|
||||
ncclInt8 = 0
|
||||
ncclChar = 0
|
||||
ncclUint8 = 1
|
||||
ncclInt32 = 2
|
||||
ncclInt = 2
|
||||
ncclUint32 = 3
|
||||
ncclInt64 = 4
|
||||
ncclUint64 = 5
|
||||
ncclFloat16 = 6
|
||||
ncclHalf = 6
|
||||
ncclFloat32 = 7
|
||||
ncclFloat = 7
|
||||
ncclFloat64 = 8
|
||||
ncclDouble = 8
|
||||
ncclBfloat16 = 9
|
||||
ncclNumTypes = 10
|
||||
|
||||
@classmethod
|
||||
def from_torch(cls, dtype: torch.dtype) -> int:
|
||||
if dtype == torch.int8:
|
||||
return cls.ncclInt8
|
||||
if dtype == torch.uint8:
|
||||
return cls.ncclUint8
|
||||
if dtype == torch.int32:
|
||||
return cls.ncclInt32
|
||||
if dtype == torch.int64:
|
||||
return cls.ncclInt64
|
||||
if dtype == torch.float16:
|
||||
return cls.ncclFloat16
|
||||
if dtype == torch.float32:
|
||||
return cls.ncclFloat32
|
||||
if dtype == torch.float64:
|
||||
return cls.ncclFloat64
|
||||
if dtype == torch.bfloat16:
|
||||
return cls.ncclBfloat16
|
||||
raise ValueError(f"Unsupported dtype: {dtype}")
|
||||
|
||||
|
||||
ncclRedOp_t = ctypes.c_int
|
||||
|
||||
|
||||
class ncclRedOpTypeEnum:
|
||||
ncclSum = 0
|
||||
ncclProd = 1
|
||||
ncclMax = 2
|
||||
ncclMin = 3
|
||||
ncclAvg = 4
|
||||
ncclNumOps = 5
|
||||
|
||||
@classmethod
|
||||
def from_torch(cls, op: ReduceOp) -> int:
|
||||
if op == ReduceOp.SUM:
|
||||
return cls.ncclSum
|
||||
if op == ReduceOp.PRODUCT:
|
||||
return cls.ncclProd
|
||||
if op == ReduceOp.MAX:
|
||||
return cls.ncclMax
|
||||
if op == ReduceOp.MIN:
|
||||
return cls.ncclMin
|
||||
if op == ReduceOp.AVG:
|
||||
return cls.ncclAvg
|
||||
raise ValueError(f"Unsupported op: {op}")
|
||||
|
||||
|
||||
@dataclass
|
||||
class Function:
|
||||
name: str
|
||||
restype: Any
|
||||
argtypes: List[Any]
|
||||
|
||||
|
||||
class NCCLLibrary:
|
||||
exported_functions = [
|
||||
# const char* ncclGetErrorString(ncclResult_t result)
|
||||
Function("ncclGetErrorString", ctypes.c_char_p, [ncclResult_t]),
|
||||
# ncclResult_t ncclGetVersion(int *version);
|
||||
Function("ncclGetVersion", ncclResult_t,
|
||||
[ctypes.POINTER(ctypes.c_int)]),
|
||||
# ncclResult_t ncclGetUniqueId(ncclUniqueId* uniqueId);
|
||||
Function("ncclGetUniqueId", ncclResult_t,
|
||||
[ctypes.POINTER(ncclUniqueId)]),
|
||||
# ncclResult_t ncclCommInitRank(
|
||||
# ncclComm_t* comm, int nranks, ncclUniqueId commId, int rank);
|
||||
# note that ncclComm_t is a pointer type, so the first argument
|
||||
# is a pointer to a pointer
|
||||
Function("ncclCommInitRank", ncclResult_t, [
|
||||
ctypes.POINTER(ncclComm_t), ctypes.c_int, ncclUniqueId,
|
||||
ctypes.c_int
|
||||
]),
|
||||
# ncclResult_t ncclAllReduce(
|
||||
# const void* sendbuff, void* recvbuff, size_t count,
|
||||
# ncclDataType_t datatype, ncclRedOp_t op, ncclComm_t comm,
|
||||
# cudaStream_t stream);
|
||||
# note that cudaStream_t is a pointer type, so the last argument
|
||||
# is a pointer
|
||||
Function("ncclAllReduce", ncclResult_t, [
|
||||
buffer_type, buffer_type, ctypes.c_size_t, ncclDataType_t,
|
||||
ncclRedOp_t, ncclComm_t, cudaStream_t
|
||||
]),
|
||||
|
||||
# ncclResult_t ncclAllGather(
|
||||
# const void* sendbuff, void* recvbuff, size_t count,
|
||||
# ncclDataType_t datatype, ncclComm_t comm,
|
||||
# cudaStream_t stream);
|
||||
# note that cudaStream_t is a pointer type, so the last argument
|
||||
# is a pointer
|
||||
Function("ncclAllGather", ncclResult_t, [
|
||||
buffer_type, buffer_type, ctypes.c_size_t, ncclDataType_t,
|
||||
ncclComm_t, cudaStream_t
|
||||
]),
|
||||
|
||||
# ncclResult_t ncclReduceScatter(
|
||||
# const void* sendbuff, void* recvbuff, size_t count,
|
||||
# ncclDataType_t datatype, ncclRedOp_t op, ncclComm_t comm,
|
||||
# cudaStream_t stream);
|
||||
# note that cudaStream_t is a pointer type, so the last argument
|
||||
# is a pointer
|
||||
Function("ncclReduceScatter", ncclResult_t, [
|
||||
buffer_type, buffer_type, ctypes.c_size_t, ncclDataType_t,
|
||||
ncclRedOp_t, ncclComm_t, cudaStream_t
|
||||
]),
|
||||
|
||||
# ncclResult_t ncclSend(
|
||||
# const void* sendbuff, size_t count, ncclDataType_t datatype,
|
||||
# int dest, ncclComm_t comm, cudaStream_t stream);
|
||||
Function("ncclSend", ncclResult_t, [
|
||||
buffer_type, ctypes.c_size_t, ncclDataType_t, ctypes.c_int,
|
||||
ncclComm_t, cudaStream_t
|
||||
]),
|
||||
|
||||
# ncclResult_t ncclRecv(
|
||||
# void* recvbuff, size_t count, ncclDataType_t datatype,
|
||||
# int src, ncclComm_t comm, cudaStream_t stream);
|
||||
Function("ncclRecv", ncclResult_t, [
|
||||
buffer_type, ctypes.c_size_t, ncclDataType_t, ctypes.c_int,
|
||||
ncclComm_t, cudaStream_t
|
||||
]),
|
||||
|
||||
# ncclResult_t ncclBroadcast(
|
||||
# const void* sendbuff, void* recvbuff, size_t count,
|
||||
# ncclDataType_t datatype, int root, ncclComm_t comm,
|
||||
# cudaStream_t stream);
|
||||
Function("ncclBroadcast", ncclResult_t, [
|
||||
buffer_type, buffer_type, ctypes.c_size_t, ncclDataType_t,
|
||||
ctypes.c_int, ncclComm_t, cudaStream_t
|
||||
]),
|
||||
|
||||
# be cautious! this is a collective call, it will block until all
|
||||
# processes in the communicator have called this function.
|
||||
# because Python object destruction can happen in random order,
|
||||
# it is better not to call it at all.
|
||||
# ncclResult_t ncclCommDestroy(ncclComm_t comm);
|
||||
Function("ncclCommDestroy", ncclResult_t, [ncclComm_t]),
|
||||
]
|
||||
|
||||
# class attribute to store the mapping from the path to the library
|
||||
# to avoid loading the same library multiple times
|
||||
path_to_library_cache: Dict[str, Any] = {}
|
||||
|
||||
# class attribute to store the mapping from library path
|
||||
# to the corresponding dictionary
|
||||
path_to_dict_mapping: Dict[str, Dict[str, Any]] = {}
|
||||
|
||||
def __init__(self, so_file: Optional[str] = None):
|
||||
|
||||
so_file = so_file or find_nccl_library()
|
||||
|
||||
try:
|
||||
if so_file not in NCCLLibrary.path_to_dict_mapping:
|
||||
lib = ctypes.CDLL(so_file)
|
||||
NCCLLibrary.path_to_library_cache[so_file] = lib
|
||||
self.lib = NCCLLibrary.path_to_library_cache[so_file]
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to load NCCL library from %s ."
|
||||
"It is expected if you are not running on NVIDIA/AMD GPUs."
|
||||
"Otherwise, the nccl library might not exist, be corrupted "
|
||||
"or it does not support the current platform %s."
|
||||
"If you already have the library, please set the "
|
||||
"environment variable VLLM_NCCL_SO_PATH"
|
||||
" to point to the correct nccl library path.", so_file,
|
||||
platform.platform())
|
||||
raise e
|
||||
|
||||
if so_file not in NCCLLibrary.path_to_dict_mapping:
|
||||
_funcs: Dict[str, Any] = {}
|
||||
for func in NCCLLibrary.exported_functions:
|
||||
f = getattr(self.lib, func.name)
|
||||
f.restype = func.restype
|
||||
f.argtypes = func.argtypes
|
||||
_funcs[func.name] = f
|
||||
NCCLLibrary.path_to_dict_mapping[so_file] = _funcs
|
||||
self._funcs = NCCLLibrary.path_to_dict_mapping[so_file]
|
||||
|
||||
def ncclGetErrorString(self, result: ncclResult_t) -> str:
|
||||
return self._funcs["ncclGetErrorString"](result).decode("utf-8")
|
||||
|
||||
def NCCL_CHECK(self, result: ncclResult_t) -> None:
|
||||
if result != 0:
|
||||
error_str = self.ncclGetErrorString(result)
|
||||
raise RuntimeError(f"NCCL error: {error_str}")
|
||||
|
||||
def ncclGetVersion(self) -> str:
|
||||
version = ctypes.c_int()
|
||||
self.NCCL_CHECK(self._funcs["ncclGetVersion"](ctypes.byref(version)))
|
||||
version_str = str(version.value)
|
||||
# something like 21903 --> "2.19.3"
|
||||
major = version_str[0].lstrip("0")
|
||||
minor = version_str[1:3].lstrip("0")
|
||||
patch = version_str[3:].lstrip("0")
|
||||
return f"{major}.{minor}.{patch}"
|
||||
|
||||
def ncclGetUniqueId(self) -> ncclUniqueId:
|
||||
unique_id = ncclUniqueId()
|
||||
self.NCCL_CHECK(self._funcs["ncclGetUniqueId"](
|
||||
ctypes.byref(unique_id)))
|
||||
return unique_id
|
||||
|
||||
def ncclCommInitRank(self, world_size: int, unique_id: ncclUniqueId,
|
||||
rank: int) -> ncclComm_t:
|
||||
comm = ncclComm_t()
|
||||
self.NCCL_CHECK(self._funcs["ncclCommInitRank"](ctypes.byref(comm),
|
||||
world_size, unique_id,
|
||||
rank))
|
||||
return comm
|
||||
|
||||
def ncclAllReduce(self, sendbuff: buffer_type, recvbuff: buffer_type,
|
||||
count: int, datatype: int, op: int, comm: ncclComm_t,
|
||||
stream: cudaStream_t) -> None:
|
||||
# `datatype` actually should be `ncclDataType_t`
|
||||
# and `op` should be `ncclRedOp_t`
|
||||
# both are aliases of `ctypes.c_int`
|
||||
# when we pass int to a function, it will be converted to `ctypes.c_int`
|
||||
# by ctypes automatically
|
||||
self.NCCL_CHECK(self._funcs["ncclAllReduce"](sendbuff, recvbuff, count,
|
||||
datatype, op, comm,
|
||||
stream))
|
||||
|
||||
def ncclReduceScatter(self, sendbuff: buffer_type, recvbuff: buffer_type,
|
||||
count: int, datatype: int, op: int, comm: ncclComm_t,
|
||||
stream: cudaStream_t) -> None:
|
||||
# `datatype` actually should be `ncclDataType_t`
|
||||
# and `op` should be `ncclRedOp_t`
|
||||
# both are aliases of `ctypes.c_int`
|
||||
# when we pass int to a function, it will be converted to `ctypes.c_int`
|
||||
# by ctypes automatically
|
||||
self.NCCL_CHECK(self._funcs["ncclReduceScatter"](sendbuff, recvbuff,
|
||||
count, datatype, op,
|
||||
comm, stream))
|
||||
|
||||
def ncclAllGather(self, sendbuff: buffer_type, recvbuff: buffer_type,
|
||||
count: int, datatype: int, comm: ncclComm_t,
|
||||
stream: cudaStream_t) -> None:
|
||||
# `datatype` actually should be `ncclDataType_t`
|
||||
# which is an aliases of `ctypes.c_int`
|
||||
# when we pass int to a function, it will be converted to `ctypes.c_int`
|
||||
# by ctypes automatically
|
||||
self.NCCL_CHECK(self._funcs["ncclAllGather"](sendbuff, recvbuff, count,
|
||||
datatype, comm, stream))
|
||||
|
||||
def ncclSend(self, sendbuff: buffer_type, count: int, datatype: int,
|
||||
dest: int, comm: ncclComm_t, stream: cudaStream_t) -> None:
|
||||
self.NCCL_CHECK(self._funcs["ncclSend"](sendbuff, count, datatype,
|
||||
dest, comm, stream))
|
||||
|
||||
def ncclRecv(self, recvbuff: buffer_type, count: int, datatype: int,
|
||||
src: int, comm: ncclComm_t, stream: cudaStream_t) -> None:
|
||||
self.NCCL_CHECK(self._funcs["ncclRecv"](recvbuff, count, datatype, src,
|
||||
comm, stream))
|
||||
|
||||
def ncclBroadcast(self, sendbuff: buffer_type, recvbuff: buffer_type,
|
||||
count: int, datatype: int, root: int, comm: ncclComm_t,
|
||||
stream: cudaStream_t) -> None:
|
||||
self.NCCL_CHECK(self._funcs["ncclBroadcast"](sendbuff, recvbuff, count,
|
||||
datatype, root, comm,
|
||||
stream))
|
||||
|
||||
def ncclCommDestroy(self, comm: ncclComm_t) -> None:
|
||||
self.NCCL_CHECK(self._funcs["ncclCommDestroy"](comm))
|
||||
|
||||
|
||||
__all__ = [
|
||||
"NCCLLibrary", "ncclDataTypeEnum", "ncclRedOpTypeEnum", "ncclUniqueId",
|
||||
"ncclComm_t", "cudaStream_t", "buffer_type"
|
||||
]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,204 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/utils.py
|
||||
|
||||
# Copyright 2023 The vLLM team.
|
||||
# Adapted from
|
||||
# https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/core/tensor_parallel/utils.py
|
||||
# Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved.
|
||||
import dataclasses
|
||||
import pickle
|
||||
import time
|
||||
from collections import deque
|
||||
from typing import Any, Deque, Dict, Optional, Sequence, Tuple
|
||||
|
||||
import torch
|
||||
from torch.distributed import ProcessGroup, TCPStore
|
||||
from torch.distributed.distributed_c10d import (Backend, PrefixStore,
|
||||
_get_default_timeout,
|
||||
is_nccl_available)
|
||||
from torch.distributed.rendezvous import rendezvous
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def ensure_divisibility(numerator, denominator):
|
||||
"""Ensure that numerator is divisible by the denominator."""
|
||||
assert numerator % denominator == 0, "{} is not divisible by {}".format(
|
||||
numerator, denominator)
|
||||
|
||||
|
||||
def divide(numerator, denominator):
|
||||
"""Ensure that numerator is divisible by the denominator and return
|
||||
the division value."""
|
||||
ensure_divisibility(numerator, denominator)
|
||||
return numerator // denominator
|
||||
|
||||
|
||||
def split_tensor_along_last_dim(
|
||||
tensor: torch.Tensor,
|
||||
num_partitions: int,
|
||||
contiguous_split_chunks: bool = False,
|
||||
) -> Sequence[torch.Tensor]:
|
||||
""" Split a tensor along its last dimension.
|
||||
|
||||
Arguments:
|
||||
tensor: input tensor.
|
||||
num_partitions: number of partitions to split the tensor
|
||||
contiguous_split_chunks: If True, make each chunk contiguous
|
||||
in memory.
|
||||
|
||||
Returns:
|
||||
A list of Tensors
|
||||
"""
|
||||
# Get the size and dimension.
|
||||
last_dim = tensor.dim() - 1
|
||||
last_dim_size = divide(tensor.size()[last_dim], num_partitions)
|
||||
# Split.
|
||||
tensor_list = torch.split(tensor, last_dim_size, dim=last_dim)
|
||||
# NOTE: torch.split does not create contiguous tensors by default.
|
||||
if contiguous_split_chunks:
|
||||
return tuple(chunk.contiguous() for chunk in tensor_list)
|
||||
|
||||
return tensor_list
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class StatelessProcessGroup:
|
||||
"""A dataclass to hold a metadata store, and the rank, world_size of the
|
||||
group. Only use it to communicate metadata between processes.
|
||||
For data-plane communication, create NCCL-related objects.
|
||||
"""
|
||||
rank: int
|
||||
world_size: int
|
||||
store: torch._C._distributed_c10d.Store
|
||||
data_expiration_seconds: int = 3600 # 1 hour
|
||||
|
||||
# dst rank -> counter
|
||||
send_dst_counter: Dict[int, int] = dataclasses.field(default_factory=dict)
|
||||
# src rank -> counter
|
||||
recv_src_counter: Dict[int, int] = dataclasses.field(default_factory=dict)
|
||||
broadcast_send_counter: int = 0
|
||||
broadcast_recv_src_counter: Dict[int, int] = dataclasses.field(
|
||||
default_factory=dict)
|
||||
|
||||
# A deque to store the data entries, with key and timestamp.
|
||||
entries: Deque[Tuple[str,
|
||||
float]] = dataclasses.field(default_factory=deque)
|
||||
|
||||
def __post_init__(self):
|
||||
assert self.rank < self.world_size
|
||||
self.send_dst_counter = {i: 0 for i in range(self.world_size)}
|
||||
self.recv_src_counter = {i: 0 for i in range(self.world_size)}
|
||||
self.broadcast_recv_src_counter = {
|
||||
i: 0
|
||||
for i in range(self.world_size)
|
||||
}
|
||||
|
||||
def send_obj(self, obj: Any, dst: int):
|
||||
"""Send an object to a destination rank."""
|
||||
self.expire_data()
|
||||
key = f"send_to/{dst}/{self.send_dst_counter[dst]}"
|
||||
self.store.set(key, pickle.dumps(obj))
|
||||
self.send_dst_counter[dst] += 1
|
||||
self.entries.append((key, time.time()))
|
||||
|
||||
def expire_data(self):
|
||||
"""Expire data that is older than `data_expiration_seconds` seconds."""
|
||||
while self.entries:
|
||||
# check the oldest entry
|
||||
key, timestamp = self.entries[0]
|
||||
if time.time() - timestamp > self.data_expiration_seconds:
|
||||
self.store.delete_key(key)
|
||||
self.entries.popleft()
|
||||
else:
|
||||
break
|
||||
|
||||
def recv_obj(self, src: int) -> Any:
|
||||
"""Receive an object from a source rank."""
|
||||
obj = pickle.loads(
|
||||
self.store.get(
|
||||
f"send_to/{self.rank}/{self.recv_src_counter[src]}"))
|
||||
self.recv_src_counter[src] += 1
|
||||
return obj
|
||||
|
||||
def broadcast_obj(self, obj: Optional[Any], src: int) -> Any:
|
||||
"""Broadcast an object from a source rank to all other ranks.
|
||||
It does not clean up after all ranks have received the object.
|
||||
Use it for limited times, e.g., for initialization.
|
||||
"""
|
||||
if self.rank == src:
|
||||
self.expire_data()
|
||||
key = (f"broadcast_from/{src}/"
|
||||
f"{self.broadcast_send_counter}")
|
||||
self.store.set(key, pickle.dumps(obj))
|
||||
self.broadcast_send_counter += 1
|
||||
self.entries.append((key, time.time()))
|
||||
return obj
|
||||
else:
|
||||
key = (f"broadcast_from/{src}/"
|
||||
f"{self.broadcast_recv_src_counter[src]}")
|
||||
recv_obj = pickle.loads(self.store.get(key))
|
||||
self.broadcast_recv_src_counter[src] += 1
|
||||
return recv_obj
|
||||
|
||||
def all_gather_obj(self, obj: Any) -> list[Any]:
|
||||
"""All gather an object from all ranks."""
|
||||
gathered_objs = []
|
||||
for i in range(self.world_size):
|
||||
if i == self.rank:
|
||||
gathered_objs.append(obj)
|
||||
self.broadcast_obj(obj, src=self.rank)
|
||||
else:
|
||||
recv_obj = self.broadcast_obj(None, src=i)
|
||||
gathered_objs.append(recv_obj)
|
||||
return gathered_objs
|
||||
|
||||
def barrier(self):
|
||||
"""A barrier to synchronize all ranks."""
|
||||
for i in range(self.world_size):
|
||||
if i == self.rank:
|
||||
self.broadcast_obj(None, src=self.rank)
|
||||
else:
|
||||
self.broadcast_obj(None, src=i)
|
||||
|
||||
@staticmethod
|
||||
def create(
|
||||
host: str,
|
||||
port: int,
|
||||
rank: int,
|
||||
world_size: int,
|
||||
data_expiration_seconds: int = 3600,
|
||||
) -> "StatelessProcessGroup":
|
||||
"""A replacement for `torch.distributed.init_process_group` that does not
|
||||
pollute the global state.
|
||||
|
||||
If we have process A and process B called `torch.distributed.init_process_group`
|
||||
to form a group, and then we want to form another group with process A, B, C,
|
||||
D, it is not possible in PyTorch, because process A and process B have already
|
||||
formed a group, and process C and process D cannot join that group. This
|
||||
function is a workaround for this issue.
|
||||
|
||||
`torch.distributed.init_process_group` is a global call, while this function
|
||||
is a stateless call. It will return a `StatelessProcessGroup` object that can be
|
||||
used for exchanging metadata. With this function, process A and process B
|
||||
can call `StatelessProcessGroup.create` to form a group, and then process A, B,
|
||||
C, and D can call `StatelessProcessGroup.create` to form another group.
|
||||
""" # noqa
|
||||
store = TCPStore(
|
||||
host_name=host,
|
||||
port=port,
|
||||
world_size=world_size,
|
||||
is_master=(rank == 0),
|
||||
)
|
||||
|
||||
return StatelessProcessGroup(
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
store=store,
|
||||
data_expiration_seconds=data_expiration_seconds)
|
||||
@@ -0,0 +1,225 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Adapted from vllm
|
||||
# https://github.com/vllm-project/vllm/blob/b382a7f28f739f3b120e5495fd029089d0399428/vllm/envs.py
|
||||
# Copyright 2023 The vLLM Authors.
|
||||
# Copyright 2023 The FastVideo Authors.
|
||||
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60
|
||||
FASTVIDEO_NCCL_SO_PATH: Optional[str] = None
|
||||
LD_LIBRARY_PATH: Optional[str] = None
|
||||
FASTVIDEO_USE_TRITON_FLASH_ATTN: bool = False
|
||||
FASTVIDEO_FLASH_ATTN_VERSION: Optional[int] = None
|
||||
LOCAL_RANK: int = 0
|
||||
CUDA_VISIBLE_DEVICES: Optional[str] = None
|
||||
FASTVIDEO_CACHE_ROOT: str = os.path.expanduser("~/.cache/fastvideo")
|
||||
FASTVIDEO_CONFIG_ROOT: str = os.path.expanduser("~/.config/fastvideo")
|
||||
FASTVIDEO_CONFIGURE_LOGGING: int = 1
|
||||
FASTVIDEO_LOGGING_LEVEL: str = "INFO"
|
||||
FASTVIDEO_LOGGING_PREFIX: str = ""
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH: Optional[str] = None
|
||||
FASTVIDEO_TRACE_FUNCTION: int = 0
|
||||
FASTVIDEO_ATTENTION_BACKEND: Optional[str] = None
|
||||
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
|
||||
FASTVIDEO_TARGET_DEVICE: str = "cuda"
|
||||
MAX_JOBS: Optional[str] = None
|
||||
NVCC_THREADS: Optional[str] = None
|
||||
CMAKE_BUILD_TYPE: Optional[str] = None
|
||||
VERBOSE: bool = False
|
||||
FASTVIDEO_SERVER_DEV_MODE: bool = False
|
||||
|
||||
|
||||
def get_default_cache_root():
|
||||
return os.getenv(
|
||||
"XDG_CACHE_HOME",
|
||||
os.path.join(os.path.expanduser("~"), ".cache"),
|
||||
)
|
||||
|
||||
|
||||
def get_default_config_root():
|
||||
return os.getenv(
|
||||
"XDG_CONFIG_HOME",
|
||||
os.path.join(os.path.expanduser("~"), ".config"),
|
||||
)
|
||||
|
||||
|
||||
def maybe_convert_int(value: Optional[str]) -> Optional[int]:
|
||||
if value is None:
|
||||
return None
|
||||
return int(value)
|
||||
|
||||
|
||||
# The begin-* and end* here are used by the documentation generator
|
||||
# to extract the used env vars.
|
||||
|
||||
# begin-env-vars-definition
|
||||
|
||||
environment_variables: Dict[str, Callable[[], Any]] = {
|
||||
|
||||
# ================== Installation Time Env Vars ==================
|
||||
|
||||
# Target device of FastVideo, supporting [cuda (by default),
|
||||
# rocm, neuron, cpu, openvino]
|
||||
"FASTVIDEO_TARGET_DEVICE":
|
||||
lambda: os.getenv("FASTVIDEO_TARGET_DEVICE", "cuda"),
|
||||
|
||||
# Maximum number of compilation jobs to run in parallel.
|
||||
# By default this is the number of CPUs
|
||||
"MAX_JOBS":
|
||||
lambda: os.getenv("MAX_JOBS", None),
|
||||
|
||||
# Number of threads to use for nvcc
|
||||
# By default this is 1.
|
||||
# If set, `MAX_JOBS` will be reduced to avoid oversubscribing the CPU.
|
||||
"NVCC_THREADS":
|
||||
lambda: os.getenv("NVCC_THREADS", None),
|
||||
|
||||
# If set, fastvideo will use precompiled binaries (*.so)
|
||||
"FASTVIDEO_USE_PRECOMPILED":
|
||||
lambda: bool(os.environ.get("FASTVIDEO_USE_PRECOMPILED")) or bool(
|
||||
os.environ.get("FASTVIDEO_PRECOMPILED_WHEEL_LOCATION")),
|
||||
|
||||
# CMake build type
|
||||
# If not set, defaults to "Debug" or "RelWithDebInfo"
|
||||
# Available options: "Debug", "Release", "RelWithDebInfo"
|
||||
"CMAKE_BUILD_TYPE":
|
||||
lambda: os.getenv("CMAKE_BUILD_TYPE"),
|
||||
|
||||
# If set, fastvideo will print verbose logs during installation
|
||||
"VERBOSE":
|
||||
lambda: bool(int(os.getenv('VERBOSE', '0'))),
|
||||
|
||||
# Root directory for FASTVIDEO configuration files
|
||||
# Defaults to `~/.config/fastvideo` unless `XDG_CONFIG_HOME` is set
|
||||
# Note that this not only affects how fastvideo finds its configuration files
|
||||
# during runtime, but also affects how fastvideo installs its configuration
|
||||
# files during **installation**.
|
||||
"FASTVIDEO_CONFIG_ROOT":
|
||||
lambda: os.path.expanduser(
|
||||
os.getenv(
|
||||
"FASTVIDEO_CONFIG_ROOT",
|
||||
os.path.join(get_default_config_root(), "fastvideo"),
|
||||
)),
|
||||
|
||||
# ================== Runtime Env Vars ==================
|
||||
|
||||
# Root directory for FASTVIDEO cache files
|
||||
# Defaults to `~/.cache/fastvideo` unless `XDG_CACHE_HOME` is set
|
||||
"FASTVIDEO_CACHE_ROOT":
|
||||
lambda: os.path.expanduser(
|
||||
os.getenv(
|
||||
"FASTVIDEO_CACHE_ROOT",
|
||||
os.path.join(get_default_cache_root(), "fastvideo"),
|
||||
)),
|
||||
|
||||
# Interval in seconds to log a warning message when the ring buffer is full
|
||||
"FASTVIDEO_RINGBUFFER_WARNING_INTERVAL":
|
||||
lambda: int(os.environ.get("FASTVIDEO_RINGBUFFER_WARNING_INTERVAL", "60")),
|
||||
|
||||
# Path to the NCCL library file. It is needed because nccl>=2.19 brought
|
||||
# by PyTorch contains a bug: https://github.com/NVIDIA/nccl/issues/1234
|
||||
"FASTVIDEO_NCCL_SO_PATH":
|
||||
lambda: os.environ.get("FASTVIDEO_NCCL_SO_PATH", None),
|
||||
|
||||
# when `FASTVIDEO_NCCL_SO_PATH` is not set, fastvideo will try to find the nccl
|
||||
# library file in the locations specified by `LD_LIBRARY_PATH`
|
||||
"LD_LIBRARY_PATH":
|
||||
lambda: os.environ.get("LD_LIBRARY_PATH", None),
|
||||
|
||||
# flag to control if fastvideo should use triton flash attention
|
||||
"FASTVIDEO_USE_TRITON_FLASH_ATTN":
|
||||
lambda: (os.environ.get("FASTVIDEO_USE_TRITON_FLASH_ATTN", "True").lower() in
|
||||
("true", "1")),
|
||||
|
||||
# Force fastvideo to use a specific flash-attention version (2 or 3), only valid
|
||||
# when using the flash-attention backend.
|
||||
"FASTVIDEO_FLASH_ATTN_VERSION":
|
||||
lambda: maybe_convert_int(os.environ.get("FASTVIDEO_FLASH_ATTN_VERSION", None)),
|
||||
|
||||
# Internal flag to enable Dynamo fullgraph capture
|
||||
"FASTVIDEO_TEST_DYNAMO_FULLGRAPH_CAPTURE":
|
||||
lambda: bool(
|
||||
os.environ.get("FASTVIDEO_TEST_DYNAMO_FULLGRAPH_CAPTURE", "1") != "0"),
|
||||
|
||||
# local rank of the process in the distributed setting, used to determine
|
||||
# the GPU device id
|
||||
"LOCAL_RANK":
|
||||
lambda: int(os.environ.get("LOCAL_RANK", "0")),
|
||||
|
||||
# used to control the visible devices in the distributed setting
|
||||
"CUDA_VISIBLE_DEVICES":
|
||||
lambda: os.environ.get("CUDA_VISIBLE_DEVICES", None),
|
||||
|
||||
# timeout for each iteration in the engine
|
||||
"FASTVIDEO_ENGINE_ITERATION_TIMEOUT_S":
|
||||
lambda: int(os.environ.get("FASTVIDEO_ENGINE_ITERATION_TIMEOUT_S", "60")),
|
||||
|
||||
# Logging configuration
|
||||
# If set to 0, fastvideo will not configure logging
|
||||
# If set to 1, fastvideo will configure logging using the default configuration
|
||||
# or the configuration file specified by FASTVIDEO_LOGGING_CONFIG_PATH
|
||||
"FASTVIDEO_CONFIGURE_LOGGING":
|
||||
lambda: int(os.getenv("FASTVIDEO_CONFIGURE_LOGGING", "1")),
|
||||
"FASTVIDEO_LOGGING_CONFIG_PATH":
|
||||
lambda: os.getenv("FASTVIDEO_LOGGING_CONFIG_PATH"),
|
||||
|
||||
# this is used for configuring the default logging level
|
||||
"FASTVIDEO_LOGGING_LEVEL":
|
||||
lambda: os.getenv("FASTVIDEO_LOGGING_LEVEL", "INFO"),
|
||||
|
||||
# if set, FASTVIDEO_LOGGING_PREFIX will be prepended to all log messages
|
||||
"FASTVIDEO_LOGGING_PREFIX":
|
||||
lambda: os.getenv("FASTVIDEO_LOGGING_PREFIX", ""),
|
||||
|
||||
# Trace function calls
|
||||
# If set to 1, fastvideo will trace function calls
|
||||
# Useful for debugging
|
||||
"FASTVIDEO_TRACE_FUNCTION":
|
||||
lambda: int(os.getenv("FASTVIDEO_TRACE_FUNCTION", "0")),
|
||||
|
||||
# Backend for attention computation
|
||||
# Available options:
|
||||
# - "TORCH_SDPA": use torch.nn.MultiheadAttention
|
||||
# - "FLASH_ATTN": use FlashAttention
|
||||
# - "XFORMERS": use XFormers
|
||||
# - "ROCM_FLASH": use ROCmFlashAttention
|
||||
# - "FLASHINFER": use flashinfer
|
||||
"FASTVIDEO_ATTENTION_BACKEND":
|
||||
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
|
||||
|
||||
# Use dedicated multiprocess context for workers.
|
||||
# Both spawn and fork work
|
||||
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
|
||||
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "fork"),
|
||||
|
||||
# Enables torch profiler if set. Path to the directory where torch profiler
|
||||
# traces are saved. Note that it must be an absolute path.
|
||||
"FASTVIDEO_TORCH_PROFILER_DIR":
|
||||
lambda: (None if os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", None) is None else os
|
||||
.path.expanduser(os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", "."))),
|
||||
|
||||
# If set, fastvideo will run in development mode, which will enable
|
||||
# some additional endpoints for developing and debugging,
|
||||
# e.g. `/reset_prefix_cache`
|
||||
"FASTVIDEO_SERVER_DEV_MODE":
|
||||
lambda: bool(int(os.getenv("FASTVIDEO_SERVER_DEV_MODE", "0"))),
|
||||
}
|
||||
|
||||
# end-env-vars-definition
|
||||
|
||||
|
||||
def __getattr__(name: str):
|
||||
# lazy evaluation of environment variables
|
||||
if name in environment_variables:
|
||||
return environment_variables[name]()
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
|
||||
def __dir__():
|
||||
return list(environment_variables.keys())
|
||||
@@ -0,0 +1,471 @@
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Adapted from SGLang server_args.py
|
||||
# Copyright 2024-2025 FastVideo Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""The arguments of FastVideo Inference."""
|
||||
|
||||
import argparse
|
||||
import dataclasses
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
from typing import List, Optional
|
||||
|
||||
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class InferenceArgs:
|
||||
# Model and path configuration
|
||||
model_path: str
|
||||
|
||||
# HuggingFace specific parameters
|
||||
trust_remote_code: bool = False
|
||||
revision: Optional[str] = None
|
||||
|
||||
# Parallelism
|
||||
tp_size: int = 1
|
||||
sp_size: int = 1
|
||||
dist_timeout: Optional[int] = None # timeout for torch.distributed
|
||||
|
||||
# Video generation parameters
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 117
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: int = 7
|
||||
|
||||
output_type: str = "pil"
|
||||
|
||||
# Model configuration
|
||||
precision: str = "bf16"
|
||||
|
||||
# VAE configurationi
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_len: int = 256
|
||||
hidden_state_skip_layer: int = 2
|
||||
|
||||
# Secondary text encoder
|
||||
text_encoder_precision_2: str = "fp16"
|
||||
text_len_2: int = 77
|
||||
|
||||
# Flow Matching parameters
|
||||
flow_solver: str = "euler"
|
||||
denoise_type: str = "flow"
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
# Scheduler options
|
||||
scheduler_type: str = "euler"
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
num_videos: int = 1
|
||||
fps: int = 24
|
||||
use_cpu_offload: bool = False
|
||||
disable_autocast: bool = False
|
||||
|
||||
|
||||
# Logging
|
||||
log_level: str = "info"
|
||||
|
||||
# Kernel backend
|
||||
attention_backend: Optional[str] = None
|
||||
|
||||
# Inference parameters
|
||||
prompt: Optional[str] = None
|
||||
prompt_path: Optional[str] = None
|
||||
output_path: str = "outputs/"
|
||||
seed: int = 1024
|
||||
device_str: Optional[str] = None
|
||||
device = None
|
||||
|
||||
def __post_init__(self):
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: argparse.ArgumentParser):
|
||||
parser.add_argument(
|
||||
"--use-v1-text-encoder",
|
||||
action="store_true",
|
||||
help="Use the v1 text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-v1-vae",
|
||||
action="store_true",
|
||||
help="Use the v1 vae",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-v1-transformer",
|
||||
action="store_true",
|
||||
help="Use the v1 transformer",
|
||||
)
|
||||
# Model and path configuration
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
type=str,
|
||||
help="The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
|
||||
required=True,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dit-weight",
|
||||
type=str,
|
||||
help="Path to the DiT model weights",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model-dir",
|
||||
type=str,
|
||||
help="Directory containing StepVideo model",
|
||||
)
|
||||
|
||||
# HuggingFace specific parameters
|
||||
parser.add_argument(
|
||||
"--trust-remote-code",
|
||||
action="store_true",
|
||||
default=InferenceArgs.trust_remote_code,
|
||||
help="Trust remote code when loading HuggingFace models",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--revision",
|
||||
type=str,
|
||||
default=InferenceArgs.revision,
|
||||
help="The specific model version to use (can be a branch name, tag name, or commit id)",
|
||||
)
|
||||
|
||||
# Parallelism
|
||||
parser.add_argument(
|
||||
"--tensor-parallel-size",
|
||||
"--tp-size",
|
||||
type=int,
|
||||
default=InferenceArgs.tp_size,
|
||||
help="The tensor parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sequence-parallel-size",
|
||||
"--sp-size",
|
||||
type=int,
|
||||
default=InferenceArgs.sp_size,
|
||||
help="The sequence parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dist-timeout",
|
||||
type=int,
|
||||
default=InferenceArgs.dist_timeout,
|
||||
help="Set timeout for torch.distributed initialization.",
|
||||
)
|
||||
|
||||
# Video generation parameters
|
||||
parser.add_argument(
|
||||
"--height",
|
||||
type=int,
|
||||
default=InferenceArgs.height,
|
||||
help="Height of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--width",
|
||||
type=int,
|
||||
default=InferenceArgs.width,
|
||||
help="Width of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-frames",
|
||||
type=int,
|
||||
default=InferenceArgs.num_frames,
|
||||
help="Number of frames to generate",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-inference-steps",
|
||||
type=int,
|
||||
default=InferenceArgs.num_inference_steps,
|
||||
help="Number of inference steps",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-scale",
|
||||
type=float,
|
||||
default=InferenceArgs.guidance_scale,
|
||||
help="Guidance scale for classifier-free guidance",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-rescale",
|
||||
type=float,
|
||||
default=InferenceArgs.guidance_rescale,
|
||||
help="Guidance rescale for classifier-free guidance",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--embedded-cfg-scale",
|
||||
type=float,
|
||||
default=InferenceArgs.embedded_cfg_scale,
|
||||
help="Embedded CFG scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flow-shift",
|
||||
"--shift",
|
||||
type=int,
|
||||
default=InferenceArgs.flow_shift,
|
||||
help="Flow shift parameter",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-type",
|
||||
type=str,
|
||||
default=InferenceArgs.output_type,
|
||||
choices=["pil"],
|
||||
help="Output type for the generated video",
|
||||
)
|
||||
|
||||
|
||||
parser.add_argument(
|
||||
"--precision",
|
||||
type=str,
|
||||
default=InferenceArgs.precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for the model",
|
||||
)
|
||||
|
||||
# VAE configuration
|
||||
parser.add_argument(
|
||||
"--vae-precision",
|
||||
type=str,
|
||||
default=InferenceArgs.vae_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for VAE",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-tiling",
|
||||
action="store_true",
|
||||
default=InferenceArgs.vae_tiling,
|
||||
help="Enable VAE tiling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-sp",
|
||||
action="store_true",
|
||||
help="Enable VAE spatial parallelism",
|
||||
)
|
||||
|
||||
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision",
|
||||
type=str,
|
||||
default=InferenceArgs.text_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-len",
|
||||
type=int,
|
||||
default=InferenceArgs.text_len,
|
||||
help="Maximum text length",
|
||||
)
|
||||
# Secondary text encoder
|
||||
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision-2",
|
||||
type=str,
|
||||
default=InferenceArgs.text_encoder_precision_2,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for secondary text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-len-2",
|
||||
type=int,
|
||||
default=InferenceArgs.text_len_2,
|
||||
help="Maximum secondary text length",
|
||||
)
|
||||
|
||||
# Flow Matching parameters
|
||||
parser.add_argument(
|
||||
"--flow-solver",
|
||||
type=str,
|
||||
default=InferenceArgs.flow_solver,
|
||||
help="Solver for flow matching",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--denoise-type",
|
||||
type=str,
|
||||
default=InferenceArgs.denoise_type,
|
||||
help="Denoise type for noised inputs",
|
||||
)
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
parser.add_argument(
|
||||
"--mask-strategy-file-path",
|
||||
type=str,
|
||||
help="Path to mask strategy JSON file for STA",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enable-torch-compile",
|
||||
action="store_true",
|
||||
help="Use torch.compile for speeding up STA inference without teacache",
|
||||
)
|
||||
|
||||
# Scheduler options
|
||||
parser.add_argument(
|
||||
"--scheduler-type",
|
||||
type=str,
|
||||
default=InferenceArgs.scheduler_type,
|
||||
help="Type of scheduler to use",
|
||||
)
|
||||
|
||||
# HunYuan specific parameters
|
||||
parser.add_argument(
|
||||
"--neg-prompt",
|
||||
type=str,
|
||||
default=InferenceArgs.neg_prompt,
|
||||
help="Negative prompt for sampling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-videos",
|
||||
type=int,
|
||||
default=InferenceArgs.num_videos,
|
||||
help="Number of videos to generate per prompt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fps",
|
||||
type=int,
|
||||
default=InferenceArgs.fps,
|
||||
help="Frames per second for output video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-cpu-offload",
|
||||
action="store_true",
|
||||
help="Use CPU offload for the model load",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-autocast",
|
||||
action="store_true",
|
||||
help="Disable autocast for denoising loop and vae decoding in pipeline sampling",
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
# Logging
|
||||
parser.add_argument(
|
||||
"--log-level",
|
||||
type=str,
|
||||
default=InferenceArgs.log_level,
|
||||
help="The logging level of all loggers.",
|
||||
)
|
||||
|
||||
# Kernel backend
|
||||
parser.add_argument(
|
||||
"--attention-backend",
|
||||
type=str,
|
||||
choices=["flashinfer", "triton", "torch_native"],
|
||||
default=InferenceArgs.attention_backend,
|
||||
help="Choose the kernels for attention layers.",
|
||||
)
|
||||
|
||||
# Inference parameters
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
type=str,
|
||||
help="Text prompt for video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt-path",
|
||||
type=str,
|
||||
help="Path to a text file containing the prompt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
type=str,
|
||||
default=InferenceArgs.output_path,
|
||||
help="Directory to save generated videos",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed",
|
||||
type=int,
|
||||
default=InferenceArgs.seed,
|
||||
help="Random seed for reproducibility",
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace):
|
||||
args.tp_size = args.tensor_parallel_size
|
||||
args.sp_size = args.sequence_parallel_size
|
||||
args.flow_shift = getattr(args, "shift", args.flow_shift)
|
||||
|
||||
# Get all fields from the dataclass
|
||||
attrs = [attr.name for attr in dataclasses.fields(cls)]
|
||||
|
||||
# Create a dictionary of attribute values, with defaults for missing attributes
|
||||
kwargs = {}
|
||||
for attr in attrs:
|
||||
# Convert snake_case attribute name to kebab-case CLI argument name
|
||||
cli_attr = attr.replace('_', '-')
|
||||
|
||||
# Handle renamed attributes or those with multiple CLI names
|
||||
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
|
||||
kwargs[attr] = args.tensor_parallel_size
|
||||
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
|
||||
kwargs[attr] = args.sequence_parallel_size
|
||||
elif attr == 'flow_shift' and hasattr(args, 'shift'):
|
||||
kwargs[attr] = args.shift
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
default_value = getattr(cls, attr, None)
|
||||
kwargs[attr] = getattr(args, attr, default_value)
|
||||
|
||||
return cls(**kwargs)
|
||||
|
||||
|
||||
def check_inference_args(self):
|
||||
"""Validate inference arguments for consistency"""
|
||||
|
||||
# Validate VAE spatial parallelism with VAE tiling
|
||||
if self.vae_sp and not self.vae_tiling:
|
||||
raise ValueError("Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True.")
|
||||
assert self.prompt is not None or self.prompt_path is not None, "Either prompt or prompt_path must be provided"
|
||||
assert self.prompt_path.endswith(".txt"), "prompt_path must be a text file"
|
||||
|
||||
_inference_args = None
|
||||
def prepare_inference_args(argv: List[str]) -> InferenceArgs:
|
||||
"""
|
||||
Prepare the inference arguments from the command line arguments.
|
||||
|
||||
Args:
|
||||
argv: The command line arguments. Typically, it should be `sys.argv[1:]`
|
||||
to ensure compatibility with `parse_args` when no arguments are passed.
|
||||
|
||||
Returns:
|
||||
The inference arguments.
|
||||
"""
|
||||
parser = FlexibleArgumentParser()
|
||||
InferenceArgs.add_cli_args(parser)
|
||||
raw_args = parser.parse_args(argv)
|
||||
inference_args = InferenceArgs.from_cli_args(raw_args)
|
||||
inference_args.check_inference_args()
|
||||
global _inference_args
|
||||
_inference_args = inference_args
|
||||
return inference_args
|
||||
|
||||
def get_inference_args() -> InferenceArgs:
|
||||
global _inference_args
|
||||
return _inference_args
|
||||
|
||||
class DeprecatedAction(argparse.Action):
|
||||
def __init__(self, option_strings, dest, nargs=0, **kwargs):
|
||||
super(DeprecatedAction, self).__init__(
|
||||
option_strings, dest, nargs=nargs, **kwargs
|
||||
)
|
||||
|
||||
def __call__(self, parser, namespace, values, option_string=None):
|
||||
raise ValueError(self.help)
|
||||
@@ -0,0 +1,213 @@
|
||||
"""
|
||||
Inference module for diffusion models.
|
||||
|
||||
This module provides classes and functions for running inference with diffusion models.
|
||||
"""
|
||||
|
||||
import os
|
||||
import time
|
||||
import torch
|
||||
from typing import Any, Dict
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.pipelines import ComposedPipelineBase, build_pipeline
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.logger import init_logger
|
||||
# TODO(will): remove, check if this is hunyuan specific
|
||||
from fastvideo.v1.utils import align_to
|
||||
# TODO(will): remove, move this to hunyuan stage
|
||||
from fastvideo.v1.pipelines.implementations.hunyuan.constants import NEGATIVE_PROMPT
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
class InferenceEngine:
|
||||
"""
|
||||
Engine for running inference with diffusion models.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pipeline: ComposedPipelineBase,
|
||||
inference_args: InferenceArgs,
|
||||
):
|
||||
"""
|
||||
Initialize the inference engine.
|
||||
|
||||
Args:
|
||||
pipeline: The pipeline to use for inference.
|
||||
inference_args: The inference arguments.
|
||||
default_negative_prompt: The default negative prompt to use.
|
||||
"""
|
||||
self.pipeline = pipeline
|
||||
self.inference_args = inference_args
|
||||
# TODO(will): this is a hack to get the default negative prompt
|
||||
self.default_negative_prompt = NEGATIVE_PROMPT
|
||||
|
||||
@classmethod
|
||||
def create_engine(
|
||||
cls,
|
||||
inference_args: InferenceArgs,
|
||||
) -> "InferenceEngine":
|
||||
"""
|
||||
Create an inference engine with the specified arguments.
|
||||
|
||||
Args:
|
||||
inference_args: The inference arguments.
|
||||
model_loader_cls: The model loader class to use. If None, it will be
|
||||
determined from the model type.
|
||||
pipeline_type: The type of pipeline to create. If None, it will be
|
||||
determined from the model type.
|
||||
|
||||
Returns:
|
||||
The created inference engine.
|
||||
|
||||
Raises:
|
||||
ValueError: If the model type is not recognized or if the pipeline type
|
||||
is not recognized.
|
||||
"""
|
||||
|
||||
logger.info(f"Building pipeline...")
|
||||
|
||||
# TODO(will): I don't really like this api.
|
||||
# it should be something closer to pipeline_cls.from_pretrained(...)
|
||||
# this way for training we can just do pipeline_cls.from_pretrained(
|
||||
# checkpoint_path) and have it handle everything.
|
||||
# TODO(Peiyuan): Then maybe we should only pass in model path and device, not the entire inference args?
|
||||
pipeline = build_pipeline(inference_args)
|
||||
logger.info(f"Pipeline Ready")
|
||||
|
||||
|
||||
# Create the inference engine
|
||||
return cls(pipeline, inference_args)
|
||||
|
||||
def run(
|
||||
self,
|
||||
prompt: str,
|
||||
inference_args: InferenceArgs,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Run inference with the pipeline.
|
||||
|
||||
Args:
|
||||
prompt: The prompt to use for generation.
|
||||
negative_prompt: The negative prompt to use. If None, the default will be used.
|
||||
seed: The random seed to use. If None, a random seed will be used.
|
||||
**kwargs: Additional arguments to pass to the pipeline.
|
||||
|
||||
Returns:
|
||||
A dictionary containing the generated videos and metadata.
|
||||
"""
|
||||
out_dict = dict()
|
||||
|
||||
num_videos_per_prompt = inference_args.num_videos
|
||||
seed = inference_args.seed
|
||||
height = inference_args.height
|
||||
width = inference_args.width
|
||||
video_length = inference_args.num_frames
|
||||
negative_prompt = inference_args.neg_prompt
|
||||
infer_steps = inference_args.num_inference_steps
|
||||
guidance_scale = inference_args.guidance_scale
|
||||
flow_shift = inference_args.flow_shift
|
||||
embedded_guidance_scale = inference_args.embedded_cfg_scale
|
||||
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# Arguments: target_width, target_height, target_video_length
|
||||
# ========================================================================
|
||||
if width <= 0 or height <= 0 or video_length <= 0:
|
||||
raise ValueError(
|
||||
f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={video_length}"
|
||||
)
|
||||
if (video_length - 1) % 4 != 0:
|
||||
raise ValueError(f"`video_length-1` must be a multiple of 4, got {video_length}")
|
||||
|
||||
logger.info(f"Input (height, width, video_length) = ({height}, {width}, {video_length})")
|
||||
|
||||
target_height = align_to(height, 16)
|
||||
target_width = align_to(width, 16)
|
||||
target_video_length = video_length
|
||||
|
||||
out_dict["size"] = (target_height, target_width, target_video_length)
|
||||
|
||||
# ========================================================================
|
||||
# Arguments: prompt, new_prompt, negative_prompt
|
||||
# ========================================================================
|
||||
if not isinstance(prompt, str):
|
||||
raise TypeError(f"`prompt` must be a string, but got {type(prompt)}")
|
||||
prompt = prompt.strip()
|
||||
|
||||
# negative prompt
|
||||
if negative_prompt is None or negative_prompt == "":
|
||||
negative_prompt = self.default_negative_prompt
|
||||
if not isinstance(negative_prompt, str):
|
||||
raise TypeError(f"`negative_prompt` must be a string, but got {type(negative_prompt)}")
|
||||
negative_prompt = negative_prompt.strip()
|
||||
|
||||
|
||||
# TODO(PY): move to hunyuan stage
|
||||
latents_size = [(video_length - 1) // 4 + 1, height // 8, width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
|
||||
# ========================================================================
|
||||
# Print infer args
|
||||
# ========================================================================
|
||||
debug_str = f"""
|
||||
height: {target_height}
|
||||
width: {target_width}
|
||||
video_length: {target_video_length}
|
||||
prompt: {prompt}
|
||||
neg_prompt: {negative_prompt}
|
||||
seed: {seed}
|
||||
infer_steps: {infer_steps}
|
||||
num_videos_per_prompt: {num_videos_per_prompt}
|
||||
guidance_scale: {guidance_scale}
|
||||
n_tokens: {n_tokens}
|
||||
flow_shift: {flow_shift}
|
||||
embedded_guidance_scale: {embedded_guidance_scale}"""
|
||||
logger.info(debug_str)
|
||||
# return
|
||||
# sp_group = get_sp_group()
|
||||
# local_rank = sp_group.rank
|
||||
device = torch.device(inference_args.device_str)
|
||||
batch = ForwardBatch(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
height=inference_args.height,
|
||||
width=inference_args.width,
|
||||
num_frames=inference_args.num_frames,
|
||||
num_inference_steps=inference_args.num_inference_steps,
|
||||
guidance_scale=inference_args.guidance_scale,
|
||||
# generator=generator,
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
data_type="video" if inference_args.num_frames > 1 else "image",
|
||||
device=device,
|
||||
extra={}, # Any additional parameters
|
||||
)
|
||||
|
||||
print('===============================================')
|
||||
print(batch)
|
||||
print('===============================================')
|
||||
print('===============================================')
|
||||
print(inference_args)
|
||||
|
||||
# ========================================================================
|
||||
# Pipeline inference
|
||||
# ========================================================================
|
||||
start_time = time.time()
|
||||
samples = self.pipeline.forward(
|
||||
batch=batch,
|
||||
inference_args=inference_args,
|
||||
)[0]
|
||||
# TODO(will): fix and move to hunyuan stage
|
||||
# out_dict["seeds"] = batch.seeds
|
||||
out_dict["samples"] = samples
|
||||
out_dict["prompts"] = prompt
|
||||
|
||||
gen_time = time.time() - start_time
|
||||
logger.info(f"Success, time: {gen_time}")
|
||||
|
||||
return out_dict
|
||||
@@ -0,0 +1,351 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Custom activation functions."""
|
||||
import math
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from vllm.distributed import (divide, get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size)
|
||||
from vllm.model_executor.custom_op import CustomOp
|
||||
from vllm.model_executor.utils import set_weight_attrs
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
|
||||
|
||||
@CustomOp.register("fatrelu_and_mul")
|
||||
class FatreluAndMul(CustomOp):
|
||||
"""An activation function for FATReLU.
|
||||
|
||||
The function computes x -> FATReLU(x[:d]) * x[d:] where
|
||||
d = x.shape[-1] // 2.
|
||||
This is used in openbmb/MiniCPM-S-1B-sft.
|
||||
|
||||
Shapes:
|
||||
x: (num_tokens, 2 * d) or (batch_size, seq_len, 2 * d)
|
||||
return: (num_tokens, d) or (batch_size, seq_len, d)
|
||||
"""
|
||||
|
||||
def __init__(self, threshold: float = 0.):
|
||||
super().__init__()
|
||||
self.threshold = threshold
|
||||
if current_platform.is_cuda_alike():
|
||||
self.op = torch.ops._C.fatrelu_and_mul
|
||||
elif current_platform.is_cpu():
|
||||
self._forward_method = self.forward_native
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
x1 = x[..., :d]
|
||||
x2 = x[..., d:]
|
||||
x1 = F.threshold(x1, self.threshold, 0.0)
|
||||
return x1 * x2
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
output_shape = (x.shape[:-1] + (d, ))
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
self.op(out, x, self.threshold)
|
||||
return out
|
||||
|
||||
|
||||
@CustomOp.register("silu_and_mul")
|
||||
class SiluAndMul(CustomOp):
|
||||
"""An activation function for SwiGLU.
|
||||
|
||||
The function computes x -> silu(x[:d]) * x[d:] where d = x.shape[-1] // 2.
|
||||
|
||||
Shapes:
|
||||
x: (num_tokens, 2 * d) or (batch_size, seq_len, 2 * d)
|
||||
return: (num_tokens, d) or (batch_size, seq_len, d)
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
if current_platform.is_cuda_alike() or current_platform.is_cpu():
|
||||
self.op = torch.ops._C.silu_and_mul
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
self.op = ipex_ops.silu_and_mul
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
d = x.shape[-1] // 2
|
||||
return F.silu(x[..., :d]) * x[..., d:]
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
output_shape = (x.shape[:-1] + (d, ))
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
output_shape = (x.shape[:-1] + (d, ))
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
|
||||
@CustomOp.register("mul_and_silu")
|
||||
class MulAndSilu(CustomOp):
|
||||
"""An activation function for SwiGLU.
|
||||
|
||||
The function computes x -> x[:d] * silu(x[d:]) where d = x.shape[-1] // 2.
|
||||
|
||||
Shapes:
|
||||
x: (num_tokens, 2 * d) or (batch_size, seq_len, 2 * d)
|
||||
return: (num_tokens, d) or (batch_size, seq_len, d)
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
if current_platform.is_cuda_alike():
|
||||
self.op = torch.ops._C.mul_and_silu
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
self.op = ipex_ops.silu_and_mul
|
||||
elif current_platform.is_cpu():
|
||||
self._forward_method = self.forward_native
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
d = x.shape[-1] // 2
|
||||
return x[..., :d] * F.silu(x[..., d:])
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
output_shape = (x.shape[:-1] + (d, ))
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
# TODO implement forward_xpu for MulAndSilu
|
||||
# def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
|
||||
|
||||
@CustomOp.register("gelu_and_mul")
|
||||
class GeluAndMul(CustomOp):
|
||||
"""An activation function for GeGLU.
|
||||
|
||||
The function computes x -> GELU(x[:d]) * x[d:] where d = x.shape[-1] // 2.
|
||||
|
||||
Shapes:
|
||||
x: (batch_size, seq_len, 2 * d) or (num_tokens, 2 * d)
|
||||
return: (batch_size, seq_len, d) or (num_tokens, d)
|
||||
"""
|
||||
|
||||
def __init__(self, approximate: str = "none"):
|
||||
super().__init__()
|
||||
self.approximate = approximate
|
||||
if approximate not in ("none", "tanh"):
|
||||
raise ValueError(f"Unknown approximate mode: {approximate}")
|
||||
if current_platform.is_cuda_alike() or current_platform.is_cpu():
|
||||
if approximate == "none":
|
||||
self.op = torch.ops._C.gelu_and_mul
|
||||
elif approximate == "tanh":
|
||||
self.op = torch.ops._C.gelu_tanh_and_mul
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
if approximate == "none":
|
||||
self.op = ipex_ops.gelu_and_mul
|
||||
else:
|
||||
self.op = ipex_ops.gelu_tanh_and_mul
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
d = x.shape[-1] // 2
|
||||
return F.gelu(x[..., :d], approximate=self.approximate) * x[..., d:]
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
output_shape = (x.shape[:-1] + (d, ))
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
output_shape = (x.shape[:-1] + (d, ))
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
def extra_repr(self) -> str:
|
||||
return f'approximate={repr(self.approximate)}'
|
||||
|
||||
|
||||
@CustomOp.register("gelu_new")
|
||||
class NewGELU(CustomOp):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
if current_platform.is_cuda_alike() or current_platform.is_cpu():
|
||||
self.op = torch.ops._C.gelu_new
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
self.op = ipex_ops.gelu_new
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
c = math.sqrt(2.0 / math.pi)
|
||||
return 0.5 * x * (1.0 + torch.tanh(c *
|
||||
(x + 0.044715 * torch.pow(x, 3.0))))
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
out = torch.empty_like(x)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.op(x)
|
||||
|
||||
|
||||
@CustomOp.register("gelu_fast")
|
||||
class FastGELU(CustomOp):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
if current_platform.is_cuda_alike() or current_platform.is_cpu():
|
||||
self.op = torch.ops._C.gelu_fast
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
self.op = ipex_ops.gelu_fast
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
return 0.5 * x * (1.0 + torch.tanh(x * 0.7978845608 *
|
||||
(1.0 + 0.044715 * x * x)))
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
out = torch.empty_like(x)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.op(x)
|
||||
|
||||
|
||||
@CustomOp.register("quick_gelu")
|
||||
class QuickGELU(CustomOp):
|
||||
# https://github.com/huggingface/transformers/blob/main/src/transformers/activations.py#L90
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
if current_platform.is_cuda_alike() or current_platform.is_cpu():
|
||||
self.op = torch.ops._C.gelu_quick
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
self.op = ipex_ops.gelu_quick
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
return x * torch.sigmoid(1.702 * x)
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
out = torch.empty_like(x)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
out = torch.empty_like(x)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
# TODO implement forward_xpu for QuickGELU
|
||||
# def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
|
||||
|
||||
@CustomOp.register("relu2")
|
||||
class ReLUSquaredActivation(CustomOp):
|
||||
"""
|
||||
Applies the relu^2 activation introduced in https://arxiv.org/abs/2109.08668v2
|
||||
"""
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
return torch.square(F.relu(x))
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.forward_native(x)
|
||||
|
||||
|
||||
class ScaledActivation(nn.Module):
|
||||
"""An activation function with post-scale parameters.
|
||||
|
||||
This is used for some quantization methods like AWQ.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
act_module: nn.Module,
|
||||
intermediate_size: int,
|
||||
input_is_parallel: bool = True,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.act = act_module
|
||||
self.input_is_parallel = input_is_parallel
|
||||
if input_is_parallel:
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
intermediate_size_per_partition = divide(intermediate_size,
|
||||
tp_size)
|
||||
else:
|
||||
intermediate_size_per_partition = intermediate_size
|
||||
if params_dtype is None:
|
||||
params_dtype = torch.get_default_dtype()
|
||||
self.scales = nn.Parameter(
|
||||
torch.empty(intermediate_size_per_partition, dtype=params_dtype))
|
||||
set_weight_attrs(self.scales, {"weight_loader": self.weight_loader})
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.act(x) / self.scales
|
||||
|
||||
def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
|
||||
param_data = param.data
|
||||
if self.input_is_parallel:
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
shard_size = param_data.shape[0]
|
||||
start_idx = tp_rank * shard_size
|
||||
loaded_weight = loaded_weight.narrow(0, start_idx, shard_size)
|
||||
assert param_data.shape == loaded_weight.shape
|
||||
param_data.copy_(loaded_weight)
|
||||
|
||||
|
||||
_ACTIVATION_REGISTRY = {
|
||||
"gelu": nn.GELU,
|
||||
"gelu_fast": FastGELU,
|
||||
"gelu_new": NewGELU,
|
||||
"gelu_pytorch_tanh": lambda: nn.GELU(approximate="tanh"),
|
||||
"relu": nn.ReLU,
|
||||
"relu2": ReLUSquaredActivation,
|
||||
"silu": nn.SiLU,
|
||||
"quick_gelu": QuickGELU,
|
||||
}
|
||||
|
||||
|
||||
def get_act_fn(act_fn_name: str) -> nn.Module:
|
||||
"""Get an activation function by name."""
|
||||
act_fn_name = act_fn_name.lower()
|
||||
if act_fn_name not in _ACTIVATION_REGISTRY:
|
||||
raise ValueError(
|
||||
f"Activation function {act_fn_name!r} is not supported.")
|
||||
|
||||
return _ACTIVATION_REGISTRY[act_fn_name]()
|
||||
|
||||
|
||||
_ACTIVATION_AND_MUL_REGISTRY = {
|
||||
"gelu": GeluAndMul,
|
||||
"silu": SiluAndMul,
|
||||
}
|
||||
|
||||
|
||||
def get_act_and_mul_fn(act_fn_name: str) -> nn.Module:
|
||||
"""Get an activation-and-mul (i.e. SiluAndMul) function by name."""
|
||||
act_fn_name = act_fn_name.lower()
|
||||
if act_fn_name not in _ACTIVATION_AND_MUL_REGISTRY:
|
||||
raise ValueError(
|
||||
f"Activation function {act_fn_name!r} is not supported.")
|
||||
|
||||
return _ACTIVATION_AND_MUL_REGISTRY[act_fn_name]()
|
||||
@@ -0,0 +1,308 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Custom normalization layers."""
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from vllm.model_executor.custom_op import CustomOp
|
||||
|
||||
|
||||
@CustomOp.register("rms_norm")
|
||||
class RMSNorm(CustomOp):
|
||||
"""Root mean square normalization.
|
||||
|
||||
Computes x -> w * x / sqrt(E[x^2] + eps) where w is the learned weight.
|
||||
Refer to https://arxiv.org/abs/1910.07467
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
eps: float = 1e-6,
|
||||
var_hidden_size: Optional[int] = None,
|
||||
has_weight: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.variance_epsilon = eps
|
||||
self.variance_size_override = (None if var_hidden_size == hidden_size
|
||||
else var_hidden_size)
|
||||
self.has_weight = has_weight
|
||||
|
||||
self.weight = torch.ones(hidden_size)
|
||||
if self.has_weight:
|
||||
self.weight = nn.Parameter(self.weight)
|
||||
|
||||
def forward_native(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
orig_dtype = x.dtype
|
||||
x = x.to(torch.float32)
|
||||
if residual is not None:
|
||||
x = x + residual.to(torch.float32)
|
||||
residual = x.to(orig_dtype)
|
||||
|
||||
hidden_size = x.shape[-1]
|
||||
if hidden_size != self.hidden_size:
|
||||
raise ValueError("Expected hidden_size to be "
|
||||
f"{self.hidden_size}, but found: {hidden_size}")
|
||||
|
||||
if self.variance_size_override is None:
|
||||
x_var = x
|
||||
else:
|
||||
if hidden_size < self.variance_size_override:
|
||||
raise ValueError(
|
||||
"Expected hidden_size to be at least "
|
||||
f"{self.variance_size_override}, but found: {hidden_size}")
|
||||
|
||||
x_var = x[:, :, :self.variance_size_override]
|
||||
|
||||
variance = x_var.pow(2).mean(dim=-1, keepdim=True)
|
||||
|
||||
x = x * torch.rsqrt(variance + self.variance_epsilon)
|
||||
x = x.to(orig_dtype)
|
||||
if self.has_weight:
|
||||
x = x * self.weight
|
||||
if residual is None:
|
||||
return x
|
||||
else:
|
||||
return x, residual
|
||||
|
||||
def forward_cuda(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
if self.variance_size_override is not None:
|
||||
return self.forward_native(x, residual)
|
||||
|
||||
from vllm import _custom_ops as ops
|
||||
|
||||
if residual is not None:
|
||||
ops.fused_add_rms_norm(
|
||||
x,
|
||||
residual,
|
||||
self.weight.data,
|
||||
self.variance_epsilon,
|
||||
)
|
||||
return x, residual
|
||||
out = torch.empty_like(x)
|
||||
ops.rms_norm(
|
||||
out,
|
||||
x,
|
||||
self.weight.data,
|
||||
self.variance_epsilon,
|
||||
)
|
||||
return out
|
||||
|
||||
def forward_hpu(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
from vllm_hpu_extension.ops import HPUFusedRMSNorm
|
||||
if HPUFusedRMSNorm is None:
|
||||
return self.forward_native(x, residual)
|
||||
if residual is not None:
|
||||
orig_shape = x.shape
|
||||
residual += x.view(residual.shape)
|
||||
# Note: HPUFusedRMSNorm requires 3D tensors as inputs
|
||||
x = HPUFusedRMSNorm.apply(residual, self.weight,
|
||||
self.variance_epsilon)
|
||||
return x.view(orig_shape), residual
|
||||
|
||||
x = HPUFusedRMSNorm.apply(x, self.weight, self.variance_epsilon)
|
||||
return x
|
||||
|
||||
def forward_xpu(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
if self.variance_size_override is not None:
|
||||
return self.forward_native(x, residual)
|
||||
|
||||
from vllm._ipex_ops import ipex_ops as ops
|
||||
|
||||
if residual is not None:
|
||||
ops.fused_add_rms_norm(
|
||||
x,
|
||||
residual,
|
||||
self.weight.data,
|
||||
self.variance_epsilon,
|
||||
)
|
||||
return x, residual
|
||||
return ops.rms_norm(
|
||||
x,
|
||||
self.weight.data,
|
||||
self.variance_epsilon,
|
||||
)
|
||||
|
||||
def extra_repr(self) -> str:
|
||||
s = f"hidden_size={self.weight.data.size(0)}"
|
||||
s += f", eps={self.variance_epsilon}"
|
||||
return s
|
||||
|
||||
|
||||
@CustomOp.register("gemma_rms_norm")
|
||||
class GemmaRMSNorm(CustomOp):
|
||||
"""RMS normalization for Gemma.
|
||||
|
||||
Two differences from the above RMSNorm:
|
||||
1. x * (1 + w) instead of x * w.
|
||||
2. (x * w).to(orig_dtype) instead of x.to(orig_dtype) * w.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
eps: float = 1e-6,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.zeros(hidden_size))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
@staticmethod
|
||||
def forward_static(
|
||||
weight: torch.Tensor,
|
||||
variance_epsilon: float,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor],
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
orig_dtype = x.dtype
|
||||
if residual is not None:
|
||||
x = x + residual
|
||||
residual = x
|
||||
|
||||
x = x.float()
|
||||
variance = x.pow(2).mean(dim=-1, keepdim=True)
|
||||
x = x * torch.rsqrt(variance + variance_epsilon)
|
||||
# Llama does x.to(float16) * w whilst Gemma is (x * w).to(float16)
|
||||
# See https://github.com/huggingface/transformers/pull/29402
|
||||
x = x * (1.0 + weight.float())
|
||||
x = x.to(orig_dtype)
|
||||
return x if residual is None else (x, residual)
|
||||
|
||||
def forward_native(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
return self.forward_static(self.weight.data, self.variance_epsilon, x,
|
||||
residual)
|
||||
|
||||
def forward_cuda(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
if torch.compiler.is_compiling():
|
||||
return self.forward_native(x, residual)
|
||||
|
||||
if not getattr(self, "_is_compiled", False):
|
||||
self.forward_static = torch.compile( # type: ignore
|
||||
self.forward_static)
|
||||
self._is_compiled = True
|
||||
return self.forward_native(x, residual)
|
||||
|
||||
|
||||
|
||||
class ScaleResidual(nn.Module):
|
||||
"""
|
||||
Applies gated residual connection.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, residual: torch.Tensor, x: torch.Tensor, gate: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply gated residual connection."""
|
||||
return residual + x * gate
|
||||
|
||||
|
||||
class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
"""
|
||||
Fused operation that combines:
|
||||
1. Gated residual connection
|
||||
2. LayerNorm
|
||||
3. Scale and shift operations
|
||||
|
||||
This reduces memory bandwidth by combining memory-bound operations.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
norm_type: str = "rms",
|
||||
eps: float = 1e-6,
|
||||
elementwise_affine: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
):
|
||||
super().__init__()
|
||||
if norm_type == "rms":
|
||||
self.norm = RMSNorm(hidden_size, has_weight=elementwise_affine, eps=eps, dtype=dtype)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=elementwise_affine, eps=eps, dtype=dtype)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
residual: torch.Tensor,
|
||||
x: torch.Tensor,
|
||||
gate: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
scale: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Apply gated residual connection, followed by layernorm and scale/shift in a single fused operation.
|
||||
|
||||
Returns:
|
||||
Tuple containing:
|
||||
- normalized and modulated output
|
||||
- residual value (value after residual connection but before normalization)
|
||||
"""
|
||||
# Apply residual connection with gating
|
||||
residual_output = residual + x * gate
|
||||
# Apply normalization
|
||||
normalized = self.norm(residual_output)
|
||||
# Apply scale and shift
|
||||
modulated = normalized * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
return modulated, residual_output
|
||||
|
||||
|
||||
|
||||
|
||||
class LayerNormScaleShift(nn.Module):
|
||||
"""
|
||||
Fused operation that combines LayerNorm with scale and shift operations.
|
||||
This reduces memory bandwidth by combining memory-bound operations.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
norm_type: str = "rms",
|
||||
eps: float = 1e-6,
|
||||
elementwise_affine: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
):
|
||||
super().__init__()
|
||||
if norm_type == "rms":
|
||||
self.norm = RMSNorm(hidden_size, has_weight=elementwise_affine, eps=eps)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=elementwise_affine, eps=eps, dtype=dtype)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
def forward(self, x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply layernorm followed by scale and shift in a single fused operation."""
|
||||
normalized = self.norm(x)
|
||||
return normalized * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,46 @@
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
"""
|
||||
MLP for DiT blocks, NO gated linear units
|
||||
TODO: add Tensor Parallel
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_dim: int,
|
||||
mlp_hidden_dim: int,
|
||||
output_dim: int = None,
|
||||
bias: bool = True,
|
||||
act_type: str = "gelu_pytorch_tanh",
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.fc_in = ReplicatedLinear(
|
||||
input_dim,
|
||||
mlp_hidden_dim, # For activation functions like SiLU that need 2x width
|
||||
bias=bias,
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
self.act = get_act_fn(act_type)
|
||||
if output_dim is None:
|
||||
output_dim = input_dim
|
||||
self.fc_out = ReplicatedLinear(
|
||||
mlp_hidden_dim,
|
||||
output_dim,
|
||||
bias=bias,
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x, _ = self.fc_in(x)
|
||||
x = self.act(x)
|
||||
x, _ = self.fc_out(x)
|
||||
return x
|
||||
|
||||
@@ -0,0 +1,541 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Adapted from
|
||||
# https://github.com/huggingface/transformers/blob/v4.33.2/src/transformers/models/llama/modeling_llama.py
|
||||
# Copyright 2023 The vLLM team.
|
||||
# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
|
||||
# and OPT implementations in this library. It has been modified from its
|
||||
# original forms to accommodate minor architectural differences compared
|
||||
# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Rotary Positional Embeddings."""
|
||||
import math
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from vllm.model_executor.custom_op import CustomOp
|
||||
from fastvideo.v1.distributed.parallel_state import get_sp_group
|
||||
|
||||
def _rotate_neox(x: torch.Tensor) -> torch.Tensor:
|
||||
x1 = x[..., :x.shape[-1] // 2]
|
||||
x2 = x[..., x.shape[-1] // 2:]
|
||||
return torch.cat((-x2, x1), dim=-1)
|
||||
|
||||
|
||||
def _rotate_gptj(x: torch.Tensor) -> torch.Tensor:
|
||||
x1 = x[..., ::2]
|
||||
x2 = x[..., 1::2]
|
||||
x = torch.stack((-x2, x1), dim=-1)
|
||||
return x.flatten(-2)
|
||||
|
||||
|
||||
def _apply_rotary_emb(
|
||||
x: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
is_neox_style: bool,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
x: [num_tokens, num_heads, head_size]
|
||||
cos: [num_tokens, head_size // 2]
|
||||
sin: [num_tokens, head_size // 2]
|
||||
is_neox_style: Whether to use the Neox-style or GPT-J-style rotary
|
||||
positional embeddings.
|
||||
"""
|
||||
# cos = cos.unsqueeze(-2).to(x.dtype)
|
||||
# sin = sin.unsqueeze(-2).to(x.dtype)
|
||||
cos = cos.unsqueeze(-2)
|
||||
sin = sin.unsqueeze(-2)
|
||||
if is_neox_style:
|
||||
x1, x2 = torch.chunk(x, 2, dim=-1)
|
||||
else:
|
||||
x1 = x[..., ::2]
|
||||
x2 = x[..., 1::2]
|
||||
o1 = (x1.float() * cos - x2.float() * sin).type_as(x)
|
||||
o2 = (x2.float() * cos + x1.float() * sin).type_as(x)
|
||||
if is_neox_style:
|
||||
return torch.cat((o1, o2), dim=-1)
|
||||
else:
|
||||
return torch.stack((o1, o2), dim=-1).flatten(-2)
|
||||
|
||||
|
||||
@CustomOp.register("rotary_embedding")
|
||||
class RotaryEmbedding(CustomOp):
|
||||
"""Original rotary positional embedding."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
head_size: int,
|
||||
rotary_dim: int,
|
||||
max_position_embeddings: int,
|
||||
base: int,
|
||||
is_neox_style: bool,
|
||||
dtype: torch.dtype,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.head_size = head_size
|
||||
self.rotary_dim = rotary_dim
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
self.base = base
|
||||
self.is_neox_style = is_neox_style
|
||||
self.dtype = dtype
|
||||
|
||||
cache = self._compute_cos_sin_cache()
|
||||
cache = cache.to(dtype)
|
||||
self.cos_sin_cache: torch.Tensor
|
||||
self.register_buffer("cos_sin_cache", cache, persistent=False)
|
||||
|
||||
def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor:
|
||||
"""Compute the inverse frequency."""
|
||||
# NOTE(woosuk): To exactly match the HF implementation, we need to
|
||||
# use CPU to compute the cache and then move it to GPU. However, we
|
||||
# create the cache on GPU for faster initialization. This may cause
|
||||
# a slight numerical difference between the HF implementation and ours.
|
||||
inv_freq = 1.0 / (base**(torch.arange(
|
||||
0, self.rotary_dim, 2, dtype=torch.float) / self.rotary_dim))
|
||||
return inv_freq
|
||||
|
||||
def _compute_cos_sin_cache(self) -> torch.Tensor:
|
||||
"""Compute the cos and sin cache."""
|
||||
inv_freq = self._compute_inv_freq(self.base)
|
||||
t = torch.arange(self.max_position_embeddings, dtype=torch.float)
|
||||
|
||||
freqs = torch.einsum("i,j -> ij", t, inv_freq)
|
||||
cos = freqs.cos()
|
||||
sin = freqs.sin()
|
||||
cache = torch.cat((cos, sin), dim=-1)
|
||||
return cache
|
||||
|
||||
def forward_native(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
offsets: Optional[torch.Tensor] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""A PyTorch-native implementation of forward()."""
|
||||
if offsets is not None:
|
||||
positions = positions + offsets
|
||||
positions = positions.flatten()
|
||||
num_tokens = positions.shape[0]
|
||||
cos_sin = self.cos_sin_cache.index_select(0, positions)
|
||||
cos, sin = cos_sin.chunk(2, dim=-1)
|
||||
|
||||
query_shape = query.shape
|
||||
query = query.view(num_tokens, -1, self.head_size)
|
||||
query_rot = query[..., :self.rotary_dim]
|
||||
query_pass = query[..., self.rotary_dim:]
|
||||
query_rot = _apply_rotary_emb(query_rot, cos, sin, self.is_neox_style)
|
||||
query = torch.cat((query_rot, query_pass), dim=-1).reshape(query_shape)
|
||||
|
||||
key_shape = key.shape
|
||||
key = key.view(num_tokens, -1, self.head_size)
|
||||
key_rot = key[..., :self.rotary_dim]
|
||||
key_pass = key[..., self.rotary_dim:]
|
||||
key_rot = _apply_rotary_emb(key_rot, cos, sin, self.is_neox_style)
|
||||
key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape)
|
||||
return query, key
|
||||
|
||||
def forward_cuda(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
offsets: Optional[torch.Tensor] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
from vllm import _custom_ops as ops
|
||||
|
||||
self.cos_sin_cache = self.cos_sin_cache.to(query.device,
|
||||
dtype=query.dtype)
|
||||
# ops.rotary_embedding()/batched_rotary_embedding()
|
||||
# are in-place operations that update the query and key tensors.
|
||||
if offsets is not None:
|
||||
ops.batched_rotary_embedding(positions, query, key, self.head_size,
|
||||
self.cos_sin_cache,
|
||||
self.is_neox_style, self.rotary_dim,
|
||||
offsets)
|
||||
else:
|
||||
ops.rotary_embedding(positions, query, key, self.head_size,
|
||||
self.cos_sin_cache, self.is_neox_style)
|
||||
return query, key
|
||||
|
||||
def forward_xpu(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
offsets: Optional[torch.Tensor] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
from vllm._ipex_ops import ipex_ops as ops
|
||||
|
||||
self.cos_sin_cache = self.cos_sin_cache.to(positions.device,
|
||||
dtype=query.dtype)
|
||||
# ops.rotary_embedding()/batched_rotary_embedding()
|
||||
# are in-place operations that update the query and key tensors.
|
||||
if offsets is not None:
|
||||
ops.batched_rotary_embedding(positions, query, key, self.head_size,
|
||||
self.cos_sin_cache,
|
||||
self.is_neox_style, self.rotary_dim,
|
||||
offsets)
|
||||
else:
|
||||
ops.rotary_embedding(positions, query, key, self.head_size,
|
||||
self.cos_sin_cache, self.is_neox_style)
|
||||
return query, key
|
||||
|
||||
def forward_hpu(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
offsets: Optional[torch.Tensor] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
from habana_frameworks.torch.hpex.kernels import (
|
||||
RotaryPosEmbeddingMode, apply_rotary_pos_emb)
|
||||
if offsets is not None:
|
||||
offsets = offsets.view(positions.shape[0], -1)
|
||||
positions = positions + offsets
|
||||
positions = positions.flatten()
|
||||
num_tokens = positions.shape[0]
|
||||
cos_sin = self.cos_sin_cache.index_select(0, positions).view(
|
||||
num_tokens, 1, -1)
|
||||
cos, sin = cos_sin.chunk(2, dim=-1)
|
||||
# HPU RoPE kernel requires hidden dimension for cos and sin to be equal
|
||||
# to query hidden dimension, so the original tensors need to be
|
||||
# expanded
|
||||
# GPT-NeoX kernel requires position_ids = None, offset, mode = BLOCKWISE
|
||||
# and expansion of cos/sin tensors via concatenation
|
||||
# GPT-J kernel requires position_ids = None, offset = 0, mode = PAIRWISE
|
||||
# and expansion of cos/sin tensors via repeat_interleave
|
||||
rope_mode: RotaryPosEmbeddingMode
|
||||
if self.is_neox_style:
|
||||
rope_mode = RotaryPosEmbeddingMode.BLOCKWISE
|
||||
cos = torch.cat((cos, cos), dim=-1)
|
||||
sin = torch.cat((sin, sin), dim=-1)
|
||||
else:
|
||||
rope_mode = RotaryPosEmbeddingMode.PAIRWISE
|
||||
sin = torch.repeat_interleave(sin,
|
||||
2,
|
||||
dim=-1,
|
||||
output_size=cos_sin.shape[-1])
|
||||
cos = torch.repeat_interleave(cos,
|
||||
2,
|
||||
dim=-1,
|
||||
output_size=cos_sin.shape[-1])
|
||||
|
||||
query_shape = query.shape
|
||||
query = query.view(num_tokens, -1, self.head_size)
|
||||
query_rot = query[..., :self.rotary_dim]
|
||||
query_pass = query[..., self.rotary_dim:]
|
||||
query_rot = apply_rotary_pos_emb(query_rot, cos, sin, None, 0,
|
||||
rope_mode)
|
||||
query = torch.cat((query_rot, query_pass), dim=-1).reshape(query_shape)
|
||||
|
||||
key_shape = key.shape
|
||||
key = key.view(num_tokens, -1, self.head_size)
|
||||
key_rot = key[..., :self.rotary_dim]
|
||||
key_pass = key[..., self.rotary_dim:]
|
||||
key_rot = apply_rotary_pos_emb(key_rot, cos, sin, None, 0, rope_mode)
|
||||
key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape)
|
||||
return query, key
|
||||
|
||||
def extra_repr(self) -> str:
|
||||
s = f"head_size={self.head_size}, rotary_dim={self.rotary_dim}"
|
||||
s += f", max_position_embeddings={self.max_position_embeddings}"
|
||||
s += f", base={self.base}, is_neox_style={self.is_neox_style}"
|
||||
return s
|
||||
|
||||
|
||||
def _to_tuple(x, dim=2):
|
||||
if isinstance(x, int):
|
||||
return (x,) * dim
|
||||
elif len(x) == dim:
|
||||
return x
|
||||
else:
|
||||
raise ValueError(f"Expected length {dim} or int, but got {x}")
|
||||
|
||||
|
||||
|
||||
def get_meshgrid_nd(start, *args, dim=2):
|
||||
"""
|
||||
Get n-D meshgrid with start, stop and num.
|
||||
|
||||
Args:
|
||||
start (int or tuple): If len(args) == 0, start is num; If len(args) == 1, start is start, args[0] is stop,
|
||||
step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num. For n-dim, start/stop/num
|
||||
should be int or n-tuple. If n-tuple is provided, the meshgrid will be stacked following the dim order in
|
||||
n-tuples.
|
||||
*args: See above.
|
||||
dim (int): Dimension of the meshgrid. Defaults to 2.
|
||||
|
||||
Returns:
|
||||
grid (np.ndarray): [dim, ...]
|
||||
"""
|
||||
if len(args) == 0:
|
||||
# start is grid_size
|
||||
num = _to_tuple(start, dim=dim)
|
||||
start = (0, ) * dim
|
||||
stop = num
|
||||
elif len(args) == 1:
|
||||
# start is start, args[0] is stop, step is 1
|
||||
start = _to_tuple(start, dim=dim)
|
||||
stop = _to_tuple(args[0], dim=dim)
|
||||
num = [stop[i] - start[i] for i in range(dim)]
|
||||
elif len(args) == 2:
|
||||
# start is start, args[0] is stop, args[1] is num
|
||||
start = _to_tuple(start, dim=dim) # Left-Top eg: 12,0
|
||||
stop = _to_tuple(args[0], dim=dim) # Right-Bottom eg: 20,32
|
||||
num = _to_tuple(args[1], dim=dim) # Target Size eg: 32,124
|
||||
else:
|
||||
raise ValueError(f"len(args) should be 0, 1 or 2, but got {len(args)}")
|
||||
|
||||
# PyTorch implement of np.linspace(start[i], stop[i], num[i], endpoint=False)
|
||||
axis_grid = []
|
||||
for i in range(dim):
|
||||
a, b, n = start[i], stop[i], num[i]
|
||||
g = torch.linspace(a, b, n + 1, dtype=torch.float32)[:n]
|
||||
axis_grid.append(g)
|
||||
grid = torch.meshgrid(*axis_grid, indexing="ij") # dim x [W, H, D]
|
||||
grid = torch.stack(grid, dim=0) # [dim, W, H, D]
|
||||
|
||||
return grid
|
||||
|
||||
|
||||
def get_1d_rotary_pos_embed(
|
||||
dim: int,
|
||||
pos: Union[torch.FloatTensor, int],
|
||||
theta: float = 10000.0,
|
||||
theta_rescale_factor: float = 1.0,
|
||||
interpolation_factor: float = 1.0,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
|
||||
(Note: `cis` means `cos + i * sin`, where i is the imaginary unit.)
|
||||
|
||||
This function calculates a frequency tensor with complex exponential using the given dimension 'dim'
|
||||
and the end index 'end'. The 'theta' parameter scales the frequencies.
|
||||
|
||||
Args:
|
||||
dim (int): Dimension of the frequency tensor.
|
||||
pos (int or torch.FloatTensor): Position indices for the frequency tensor. [S] or scalar
|
||||
theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.
|
||||
theta_rescale_factor (float, optional): Rescale factor for theta. Defaults to 1.0.
|
||||
interpolation_factor (float, optional): Factor to scale positions. Defaults to 1.0.
|
||||
|
||||
Returns:
|
||||
freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately. [S, D]
|
||||
"""
|
||||
if isinstance(pos, int):
|
||||
pos = torch.arange(pos).float()
|
||||
|
||||
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
|
||||
# has some connection to NTK literature
|
||||
if theta_rescale_factor != 1.0:
|
||||
theta *= theta_rescale_factor**(dim / (dim - 2))
|
||||
|
||||
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].to(torch.float64) / dim)) # [D/2]
|
||||
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
|
||||
freqs_cos = freqs.cos() # [S, D/2]
|
||||
freqs_sin = freqs.sin() # [S, D/2]
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
def get_nd_rotary_pos_embed(
|
||||
rope_dim_list,
|
||||
start,
|
||||
*args,
|
||||
theta=10000.0,
|
||||
theta_rescale_factor: Union[float, List[float]] = 1.0,
|
||||
interpolation_factor: Union[float, List[float]] = 1.0,
|
||||
shard_dim: int = 0,
|
||||
sp_rank: int = 0,
|
||||
sp_world_size: int = 1,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
|
||||
Supports sequence parallelism by allowing sharding of a specific dimension.
|
||||
|
||||
Args:
|
||||
rope_dim_list (list of int): Dimension of each rope. len(rope_dim_list) should equal to n.
|
||||
sum(rope_dim_list) should equal to head_dim of attention layer.
|
||||
start (int | tuple of int | list of int): If len(args) == 0, start is num; If len(args) == 1, start is start,
|
||||
args[0] is stop, step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num.
|
||||
*args: See above.
|
||||
theta (float): Scaling factor for frequency computation. Defaults to 10000.0.
|
||||
theta_rescale_factor (float): Rescale factor for theta. Defaults to 1.0.
|
||||
interpolation_factor (float): Factor to scale positions. Defaults to 1.0.
|
||||
shard_dim (int): Which dimension to shard for sequence parallelism. Defaults to 0.
|
||||
sp_rank (int): Rank in the sequence parallel group. Defaults to 0.
|
||||
sp_world_size (int): World size of the sequence parallel group. Defaults to 1.
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, torch.Tensor]: (cos, sin) tensors of shape [HW, D/2]
|
||||
"""
|
||||
# Get the full grid
|
||||
full_grid = get_meshgrid_nd(start, *args, dim=len(rope_dim_list)) # [3, W, H, D] / [2, W, H]
|
||||
|
||||
# Shard the grid if using sequence parallelism (sp_world_size > 1)
|
||||
assert shard_dim < len(rope_dim_list), f"shard_dim {shard_dim} must be less than number of dimensions {len(rope_dim_list)}"
|
||||
if sp_world_size > 1:
|
||||
# Get the shape of the full grid
|
||||
grid_shape = list(full_grid.shape[1:])
|
||||
|
||||
# Ensure the dimension to shard is divisible by sp_world_size
|
||||
assert grid_shape[shard_dim] % sp_world_size == 0, (
|
||||
f"Dimension {shard_dim} with size {grid_shape[shard_dim]} is not divisible "
|
||||
f"by sequence parallel world size {sp_world_size}"
|
||||
)
|
||||
|
||||
# Compute the start and end indices for this rank's shard
|
||||
shard_size = grid_shape[shard_dim] // sp_world_size
|
||||
start_idx = sp_rank * shard_size
|
||||
end_idx = (sp_rank + 1) * shard_size
|
||||
|
||||
# Create slicing indices for each dimension
|
||||
slice_indices = [slice(None) for _ in range(len(grid_shape))]
|
||||
slice_indices[shard_dim] = slice(start_idx, end_idx)
|
||||
|
||||
# Shard the grid
|
||||
# Update grid shape for the sharded dimension
|
||||
grid_shape[shard_dim] = grid_shape[shard_dim] // sp_world_size
|
||||
grid = torch.empty((len(rope_dim_list),) + tuple(grid_shape), dtype=full_grid.dtype)
|
||||
for i in range(len(rope_dim_list)):
|
||||
grid[i] = full_grid[i][tuple(slice_indices)]
|
||||
else:
|
||||
grid = full_grid
|
||||
|
||||
if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):
|
||||
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
|
||||
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
|
||||
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
|
||||
assert len(theta_rescale_factor) == len(
|
||||
rope_dim_list), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
|
||||
|
||||
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
|
||||
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
|
||||
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
|
||||
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
|
||||
assert len(interpolation_factor) == len(
|
||||
rope_dim_list), "len(interpolation_factor) should equal to len(rope_dim_list)"
|
||||
|
||||
# use 1/ndim of dimensions to encode grid_axis
|
||||
embs = []
|
||||
for i in range(len(rope_dim_list)):
|
||||
emb = get_1d_rotary_pos_embed(
|
||||
rope_dim_list[i],
|
||||
grid[i].reshape(-1),
|
||||
theta,
|
||||
theta_rescale_factor=theta_rescale_factor[i],
|
||||
interpolation_factor=interpolation_factor[i],
|
||||
) # 2 x [WHD, rope_dim_list[i]]
|
||||
embs.append(emb)
|
||||
|
||||
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D/2)
|
||||
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D/2)
|
||||
return cos, sin
|
||||
|
||||
|
||||
def get_rotary_pos_embed(
|
||||
rope_sizes,
|
||||
hidden_size,
|
||||
heads_num,
|
||||
rope_dim_list,
|
||||
rope_theta,
|
||||
theta_rescale_factor=1.0,
|
||||
interpolation_factor=1.0,
|
||||
shard_dim: int = 0,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Generate rotary positional embeddings for the given sizes.
|
||||
|
||||
Args:
|
||||
rope_sizes: Tuple of dimensions (t, h, w)
|
||||
hidden_size: Hidden dimension size
|
||||
heads_num: Number of attention heads
|
||||
rope_dim_list: List of dimensions for each axis, or None
|
||||
rope_theta: Base for frequency calculations
|
||||
theta_rescale_factor: Rescale factor for theta. Defaults to 1.0
|
||||
interpolation_factor: Factor to scale positions. Defaults to 1.0
|
||||
shard_dim: Which dimension to shard for sequence parallelism. Defaults to 0.
|
||||
|
||||
Returns:
|
||||
Tuple of (cos, sin) tensors for rotary embeddings
|
||||
"""
|
||||
|
||||
target_ndim = 3
|
||||
head_dim = hidden_size // heads_num
|
||||
|
||||
if rope_dim_list is None:
|
||||
rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)]
|
||||
|
||||
assert sum(rope_dim_list) == head_dim, "sum(rope_dim_list) should equal to head_dim of attention layer"
|
||||
|
||||
# Get SP info
|
||||
sp_group = get_sp_group()
|
||||
sp_rank = sp_group.rank_in_group
|
||||
sp_world_size = sp_group.world_size
|
||||
|
||||
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
|
||||
rope_dim_list,
|
||||
rope_sizes,
|
||||
theta=rope_theta,
|
||||
theta_rescale_factor=theta_rescale_factor,
|
||||
interpolation_factor=interpolation_factor,
|
||||
shard_dim=shard_dim,
|
||||
sp_rank=sp_rank,
|
||||
sp_world_size=sp_world_size
|
||||
)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
_ROPE_DICT: Dict[Tuple, RotaryEmbedding] = {}
|
||||
|
||||
def get_rope(
|
||||
head_size: int,
|
||||
rotary_dim: int,
|
||||
max_position: int,
|
||||
base: int,
|
||||
is_neox_style: bool = True,
|
||||
rope_scaling: Optional[Dict[str, Any]] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
partial_rotary_factor: float = 1.0,
|
||||
) -> RotaryEmbedding:
|
||||
if dtype is None:
|
||||
dtype = torch.get_default_dtype()
|
||||
if rope_scaling is not None:
|
||||
# Transforms every value that is a list into a tuple for caching calls
|
||||
rope_scaling_tuple = {
|
||||
k: tuple(v) if isinstance(v, list) else v
|
||||
for k, v in rope_scaling.items()
|
||||
}
|
||||
rope_scaling_args = tuple(rope_scaling_tuple.items())
|
||||
else:
|
||||
rope_scaling_args = None
|
||||
if partial_rotary_factor < 1.0:
|
||||
rotary_dim = int(rotary_dim * partial_rotary_factor)
|
||||
key = (head_size, rotary_dim, max_position, base, is_neox_style,
|
||||
rope_scaling_args, dtype)
|
||||
if key in _ROPE_DICT:
|
||||
return _ROPE_DICT[key]
|
||||
|
||||
if rope_scaling is None:
|
||||
rotary_emb = RotaryEmbedding(head_size, rotary_dim, max_position, base,
|
||||
is_neox_style, dtype)
|
||||
else:
|
||||
raise ValueError(f"Unknown RoPE scaling {rope_scaling}")
|
||||
_ROPE_DICT[key] = rotary_emb
|
||||
return rotary_emb
|
||||
@@ -0,0 +1,58 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Utility methods for model layers."""
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def get_token_bin_counts_and_mask(
|
||||
tokens: torch.Tensor,
|
||||
vocab_size: int,
|
||||
num_seqs: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# Compute the bin counts for the tokens.
|
||||
# vocab_size + 1 for padding.
|
||||
bin_counts = torch.zeros((num_seqs, vocab_size + 1),
|
||||
dtype=torch.long,
|
||||
device=tokens.device)
|
||||
bin_counts.scatter_add_(1, tokens, torch.ones_like(tokens))
|
||||
bin_counts = bin_counts[:, :vocab_size]
|
||||
mask = bin_counts > 0
|
||||
|
||||
return bin_counts, mask
|
||||
|
||||
|
||||
def apply_penalties(logits: torch.Tensor, prompt_tokens_tensor: torch.Tensor,
|
||||
output_tokens_tensor: torch.Tensor,
|
||||
presence_penalties: torch.Tensor,
|
||||
frequency_penalties: torch.Tensor,
|
||||
repetition_penalties: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Applies penalties in place to the logits tensor
|
||||
logits : The input logits tensor of shape [num_seqs, vocab_size]
|
||||
prompt_tokens_tensor: A tensor containing the prompt tokens. The prompts
|
||||
are padded to the maximum prompt length within the batch using
|
||||
`vocab_size` as the padding value. The value `vocab_size` is used
|
||||
for padding because it does not correspond to any valid token ID
|
||||
in the vocabulary.
|
||||
output_tokens_tensor: The output tokens tensor.
|
||||
presence_penalties: The presence penalties of shape (num_seqs, )
|
||||
frequency_penalties: The frequency penalties of shape (num_seqs, )
|
||||
repetition_penalties: The repetition penalties of shape (num_seqs, )
|
||||
"""
|
||||
num_seqs, vocab_size = logits.shape
|
||||
_, prompt_mask = get_token_bin_counts_and_mask(prompt_tokens_tensor,
|
||||
vocab_size, num_seqs)
|
||||
output_bin_counts, output_mask = get_token_bin_counts_and_mask(
|
||||
output_tokens_tensor, vocab_size, num_seqs)
|
||||
repetition_penalties = repetition_penalties.unsqueeze(dim=1).repeat(
|
||||
1, vocab_size)
|
||||
logits[logits > 0] /= torch.where(prompt_mask | output_mask,
|
||||
repetition_penalties, 1.0)[logits > 0]
|
||||
logits[logits <= 0] *= torch.where(prompt_mask | output_mask,
|
||||
repetition_penalties, 1.0)[logits <= 0]
|
||||
# We follow the definition in OpenAI API.
|
||||
# Refer to https://platform.openai.com/docs/api-reference/parameter-details
|
||||
logits -= frequency_penalties.unsqueeze(dim=1) * output_bin_counts
|
||||
logits -= presence_penalties.unsqueeze(dim=1) * output_mask
|
||||
return logits
|
||||
@@ -0,0 +1,165 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import math
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
from typing import Optional
|
||||
from fastvideo.v1.layers.mlp import MLP
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
"""2D Image to Patch Embedding
|
||||
|
||||
Image to Patch Embedding using Conv2d
|
||||
|
||||
A convolution based approach to patchifying a 2D image w/ embedding projection.
|
||||
|
||||
Based on the impl in https://github.com/google-research/vision_transformer
|
||||
|
||||
Hacked together by / Copyright 2020 Ross Wightman
|
||||
|
||||
Remove the _assert function in forward function to be compatible with multi-resolution images.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=16,
|
||||
in_chans=3,
|
||||
embed_dim=768,
|
||||
norm_layer=None,
|
||||
flatten=True,
|
||||
bias=True,
|
||||
dtype=None
|
||||
):
|
||||
super().__init__()
|
||||
# Convert patch_size to 2-tuple
|
||||
if isinstance(patch_size, (list, tuple)):
|
||||
if len(patch_size) == 1:
|
||||
patch_size = (patch_size[0], patch_size[0])
|
||||
else:
|
||||
patch_size = (patch_size, patch_size)
|
||||
|
||||
self.patch_size = patch_size
|
||||
self.flatten = flatten
|
||||
|
||||
self.proj = nn.Conv3d(
|
||||
in_chans,
|
||||
embed_dim,
|
||||
kernel_size=patch_size,
|
||||
stride=patch_size,
|
||||
bias=bias,
|
||||
dtype=dtype
|
||||
)
|
||||
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.proj(x)
|
||||
if self.flatten:
|
||||
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
|
||||
x = self.norm(x)
|
||||
return x
|
||||
|
||||
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
act_layer="silu",
|
||||
frequency_embedding_size=256,
|
||||
max_period=10000,
|
||||
dtype=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
self.max_period = max_period
|
||||
|
||||
self.mlp = MLP(
|
||||
frequency_embedding_size,
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
act_type=act_layer,
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
def forward(self, t):
|
||||
t_freq = timestep_embedding(t, self.frequency_embedding_size, self.max_period).float()
|
||||
# t_freq = t_freq.to(self.mlp.fc_in.weight.dtype)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
|
||||
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
|
||||
Args:
|
||||
t: Tensor of shape [B] with timesteps
|
||||
dim: Embedding dimension
|
||||
max_period: Controls the minimum frequency of the embeddings
|
||||
|
||||
Returns:
|
||||
Tensor of shape [B, dim] with embeddings
|
||||
"""
|
||||
half = dim // 2
|
||||
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float64) / half).to(device=t.device)
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
return embedding
|
||||
|
||||
|
||||
class ModulateProjection(nn.Module):
|
||||
"""Modulation layer for DiT blocks."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
factor: int = 2,
|
||||
act_layer: str = "silu",
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.factor = factor
|
||||
self.hidden_size = hidden_size
|
||||
self.linear = ReplicatedLinear(
|
||||
hidden_size,
|
||||
hidden_size * factor,
|
||||
bias=True,
|
||||
params_dtype=dtype
|
||||
)
|
||||
self.act = get_act_fn(act_layer)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.act(x)
|
||||
x, _ = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
def unpatchify(x, t, h, w, patch_size, channels):
|
||||
"""
|
||||
Convert patched representation back to image space.
|
||||
|
||||
Args:
|
||||
x: Tensor of shape [B, T*H*W, C*P_t*P_h*P_w]
|
||||
t, h, w: Temporal and spatial dimensions
|
||||
|
||||
Returns:
|
||||
Unpatchified tensor of shape [B, C, T*P_t, H*P_h, W*P_w]
|
||||
"""
|
||||
assert x.ndim == 3, f"x.ndim: {x.ndim}"
|
||||
assert len(patch_size) == 3, f"patch_size: {patch_size}"
|
||||
assert t * h * w == x.shape[1], f"t * h * w: {t * h * w}, x.shape[1]: {x.shape[1]}"
|
||||
c = channels
|
||||
pt, ph, pw = patch_size
|
||||
|
||||
x = x.reshape(shape=(x.shape[0], t, h, w, c, pt, ph, pw))
|
||||
x = torch.einsum("nthwcopq->nctohpwq", x)
|
||||
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
|
||||
|
||||
return imgs
|
||||
@@ -0,0 +1,484 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Sequence, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.nn.parameter import Parameter, UninitializedParameter
|
||||
|
||||
from fastvideo.v1.distributed import (divide, get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size,
|
||||
tensor_model_parallel_all_reduce)
|
||||
from vllm.model_executor.layers.quantization.base_config import (
|
||||
QuantizationConfig, QuantizeMethodBase, method_has_implemented_embedding)
|
||||
from fastvideo.v1.models.parameter import BasevLLMParameter
|
||||
from fastvideo.v1.models.utils import set_weight_attrs
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
|
||||
DEFAULT_VOCAB_PADDING_SIZE = 64
|
||||
|
||||
|
||||
class UnquantizedEmbeddingMethod(QuantizeMethodBase):
|
||||
"""Unquantized method for embeddings."""
|
||||
|
||||
def create_weights(self, layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: List[int], input_size: int,
|
||||
output_size: int, params_dtype: torch.dtype,
|
||||
**extra_weight_attrs):
|
||||
"""Create weights for embedding layer."""
|
||||
weight = Parameter(torch.empty(sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype),
|
||||
requires_grad=False)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
def apply(self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
return F.linear(x, layer.weight, bias)
|
||||
|
||||
def embedding(self, layer: torch.nn.Module,
|
||||
input_: torch.Tensor) -> torch.Tensor:
|
||||
return F.embedding(input_, layer.weight)
|
||||
|
||||
|
||||
def pad_vocab_size(vocab_size: int,
|
||||
pad_to: int = DEFAULT_VOCAB_PADDING_SIZE) -> int:
|
||||
"""Pad the vocab size to the given value."""
|
||||
return ((vocab_size + pad_to - 1) // pad_to) * pad_to
|
||||
|
||||
|
||||
def vocab_range_from_per_partition_vocab_size(
|
||||
per_partition_vocab_size: int,
|
||||
rank: int,
|
||||
offset: int = 0) -> Sequence[int]:
|
||||
index_f = rank * per_partition_vocab_size
|
||||
index_l = index_f + per_partition_vocab_size
|
||||
return index_f + offset, index_l + offset
|
||||
|
||||
|
||||
def vocab_range_from_global_vocab_size(global_vocab_size: int,
|
||||
rank: int,
|
||||
world_size: int,
|
||||
offset: int = 0) -> Sequence[int]:
|
||||
per_partition_vocab_size = divide(global_vocab_size, world_size)
|
||||
return vocab_range_from_per_partition_vocab_size(per_partition_vocab_size,
|
||||
rank,
|
||||
offset=offset)
|
||||
|
||||
|
||||
@dataclass
|
||||
class VocabParallelEmbeddingShardIndices:
|
||||
"""Indices for a shard of a vocab parallel embedding."""
|
||||
padded_org_vocab_start_index: int
|
||||
padded_org_vocab_end_index: int
|
||||
padded_added_vocab_start_index: int
|
||||
padded_added_vocab_end_index: int
|
||||
|
||||
org_vocab_start_index: int
|
||||
org_vocab_end_index: int
|
||||
added_vocab_start_index: int
|
||||
added_vocab_end_index: int
|
||||
|
||||
@property
|
||||
def num_org_elements(self) -> int:
|
||||
return self.org_vocab_end_index - self.org_vocab_start_index
|
||||
|
||||
@property
|
||||
def num_added_elements(self) -> int:
|
||||
return self.added_vocab_end_index - self.added_vocab_start_index
|
||||
|
||||
@property
|
||||
def num_org_elements_padded(self) -> int:
|
||||
return (self.padded_org_vocab_end_index -
|
||||
self.padded_org_vocab_start_index)
|
||||
|
||||
@property
|
||||
def num_added_elements_padded(self) -> int:
|
||||
return (self.padded_added_vocab_end_index -
|
||||
self.padded_added_vocab_start_index)
|
||||
|
||||
@property
|
||||
def num_org_vocab_padding(self) -> int:
|
||||
return self.num_org_elements_padded - self.num_org_elements
|
||||
|
||||
@property
|
||||
def num_added_vocab_padding(self) -> int:
|
||||
return self.num_added_elements_padded - self.num_added_elements
|
||||
|
||||
@property
|
||||
def num_elements_padded(self) -> int:
|
||||
return self.num_org_elements_padded + self.num_added_elements_padded
|
||||
|
||||
def __post_init__(self):
|
||||
# sanity checks
|
||||
assert (self.padded_org_vocab_start_index
|
||||
<= self.padded_org_vocab_end_index)
|
||||
assert (self.padded_added_vocab_start_index
|
||||
<= self.padded_added_vocab_end_index)
|
||||
|
||||
assert self.org_vocab_start_index <= self.org_vocab_end_index
|
||||
assert self.added_vocab_start_index <= self.added_vocab_end_index
|
||||
|
||||
assert self.org_vocab_start_index <= self.padded_org_vocab_start_index
|
||||
assert (self.added_vocab_start_index
|
||||
<= self.padded_added_vocab_start_index)
|
||||
assert self.org_vocab_end_index <= self.padded_org_vocab_end_index
|
||||
assert self.added_vocab_end_index <= self.padded_added_vocab_end_index
|
||||
|
||||
assert self.num_org_elements <= self.num_org_elements_padded
|
||||
assert self.num_added_elements <= self.num_added_elements_padded
|
||||
|
||||
|
||||
@torch.compile(dynamic=True, backend=current_platform.simple_compile_backend)
|
||||
def get_masked_input_and_mask(
|
||||
input_: torch.Tensor, org_vocab_start_index: int,
|
||||
org_vocab_end_index: int, num_org_vocab_padding: int,
|
||||
added_vocab_start_index: int,
|
||||
added_vocab_end_index: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# torch.compile will fuse all of the pointwise ops below
|
||||
# into a single kernel, making it very fast
|
||||
org_vocab_mask = (input_ >= org_vocab_start_index) & (
|
||||
input_ < org_vocab_end_index)
|
||||
added_vocab_mask = (input_ >= added_vocab_start_index) & (
|
||||
input_ < added_vocab_end_index)
|
||||
added_offset = added_vocab_start_index - (
|
||||
org_vocab_end_index - org_vocab_start_index) - num_org_vocab_padding
|
||||
valid_offset = (org_vocab_start_index *
|
||||
org_vocab_mask) + (added_offset * added_vocab_mask)
|
||||
vocab_mask = org_vocab_mask | added_vocab_mask
|
||||
input_ = vocab_mask * (input_ - valid_offset)
|
||||
return input_, ~vocab_mask
|
||||
|
||||
|
||||
class VocabParallelEmbedding(torch.nn.Module):
|
||||
"""Embedding parallelized in the vocabulary dimension.
|
||||
|
||||
Adapted from torch.nn.Embedding, note that we pad the vocabulary size to
|
||||
make sure it is divisible by the number of model parallel GPUs.
|
||||
|
||||
In order to support various loading methods, we ensure that LoRA-added
|
||||
embeddings are always at the end of TP-sharded tensors. In other words,
|
||||
we shard base embeddings and LoRA embeddings separately (both padded),
|
||||
and place them in the same tensor.
|
||||
In this example, we will have the original vocab size = 1010,
|
||||
added vocab size = 16 and padding to 64. Therefore, the total
|
||||
vocab size with padding will be 1088 (because we first pad 1010 to
|
||||
1024, add 16, and then pad to 1088).
|
||||
Therefore, the tensor format looks like the following:
|
||||
TP1, rank 0 (no sharding):
|
||||
|< --------BASE-------- >|< -BASE PADDING-- >|< -----LORA------ >|< -LORA PADDING-- >|
|
||||
corresponding token_id: | 0 | 1 | ... | 1009 | -1 | ... | -1 | 1010 | ... | 1015 | -1 | ... | -1 |
|
||||
index: | 0 | 1 | ... | 1009 | 1010 | ... | 1023 | 1024 | ... | 1039 | 1040 | ... | 1087 |
|
||||
|
||||
TP2, rank 0:
|
||||
|< --------------------BASE--------------------- >|< -----LORA------ >|< -LORA PADDING- >|
|
||||
corresponding token_id: | 0 | 1 | 2 | ... | 497 | 498 | ... | 511 | 1000 | ... | 1015 | -1 | ... | -1 |
|
||||
index: | 0 | 1 | 2 | ... | 497 | 498 | ... | 511 | 512 | ... | 527 | 520 | ... | 543 |
|
||||
TP2, rank 1:
|
||||
|< -----------BASE----------- >|< -BASE PADDING- >|< -----------LORA PADDING----------- >|
|
||||
corresponding token_id: | 512 | 513 | 514 | ... | 1009 | -1 | ... | -1 | -1 | ... | -1 | -1 | ... | -1 |
|
||||
index: | 0 | 1 | 2 | ... | 497 | 498 | ... | 511 | 512 | ... | 519 | 520 | ... | 543 |
|
||||
|
||||
Args:
|
||||
num_embeddings: vocabulary size.
|
||||
embedding_dim: size of hidden state.
|
||||
params_dtype: type of the parameters.
|
||||
org_num_embeddings: original vocabulary size (without LoRA).
|
||||
padding_size: padding size for the vocabulary.
|
||||
quant_config: quant config for the layer
|
||||
prefix: full name of the layer in the state dict
|
||||
""" # noqa: E501
|
||||
|
||||
def __init__(self,
|
||||
num_embeddings: int,
|
||||
embedding_dim: int,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
org_num_embeddings: Optional[int] = None,
|
||||
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
# Keep the input dimensions.
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
self.num_embeddings = num_embeddings
|
||||
self.padding_size = padding_size
|
||||
self.org_vocab_size = org_num_embeddings or num_embeddings
|
||||
num_added_embeddings = num_embeddings - self.org_vocab_size
|
||||
self.org_vocab_size_padded = pad_vocab_size(self.org_vocab_size,
|
||||
self.padding_size)
|
||||
self.num_embeddings_padded = pad_vocab_size(
|
||||
self.org_vocab_size_padded + num_added_embeddings,
|
||||
self.padding_size)
|
||||
assert self.org_vocab_size_padded <= self.num_embeddings_padded
|
||||
|
||||
self.shard_indices = self._get_indices(self.num_embeddings_padded,
|
||||
self.org_vocab_size_padded,
|
||||
self.num_embeddings,
|
||||
self.org_vocab_size, tp_rank,
|
||||
self.tp_size)
|
||||
self.embedding_dim = embedding_dim
|
||||
|
||||
quant_method = None
|
||||
if quant_config is not None:
|
||||
quant_method = quant_config.get_quant_method(self, prefix=prefix)
|
||||
if quant_method is None:
|
||||
quant_method = UnquantizedEmbeddingMethod()
|
||||
|
||||
# If we are making an embedding layer, then our quantization linear
|
||||
# method must implement the embedding operation. If we are another
|
||||
# layer type like ParallelLMHead, this is not important.
|
||||
is_embedding_layer = type(self.__class__) is VocabParallelEmbedding
|
||||
quant_method_implements_embedding = method_has_implemented_embedding(
|
||||
type(quant_method))
|
||||
if is_embedding_layer and not quant_method_implements_embedding:
|
||||
raise NotImplementedError(
|
||||
f"The class {type(quant_method).__name__} must implement "
|
||||
"the 'embedding' method, see UnquantizedEmbeddingMethod.")
|
||||
|
||||
self.quant_method: QuantizeMethodBase = quant_method
|
||||
|
||||
if params_dtype is None:
|
||||
params_dtype = torch.get_default_dtype()
|
||||
# Divide the weight matrix along the vocaburaly dimension.
|
||||
self.num_added_embeddings = self.num_embeddings - self.org_vocab_size
|
||||
self.num_embeddings_per_partition = divide(self.num_embeddings_padded,
|
||||
self.tp_size)
|
||||
assert (self.shard_indices.num_elements_padded ==
|
||||
self.num_embeddings_per_partition)
|
||||
self.num_org_embeddings_per_partition = (
|
||||
self.shard_indices.org_vocab_end_index -
|
||||
self.shard_indices.org_vocab_start_index)
|
||||
self.num_added_embeddings_per_partition = (
|
||||
self.shard_indices.added_vocab_end_index -
|
||||
self.shard_indices.added_vocab_start_index)
|
||||
|
||||
self.quant_method.create_weights(self,
|
||||
self.embedding_dim,
|
||||
[self.num_embeddings_per_partition],
|
||||
self.embedding_dim,
|
||||
self.num_embeddings_padded,
|
||||
params_dtype=params_dtype,
|
||||
weight_loader=self.weight_loader)
|
||||
|
||||
@classmethod
|
||||
def _get_indices(cls, vocab_size_padded: int, org_vocab_size_padded: int,
|
||||
vocab_size: int, org_vocab_size: int, tp_rank: int,
|
||||
tp_size: int) -> VocabParallelEmbeddingShardIndices:
|
||||
"""Get start and end indices for vocab parallel embedding, following the
|
||||
layout outlined in the class docstring, based on the given tp_rank and
|
||||
tp_size."""
|
||||
num_added_embeddings_padded = vocab_size_padded - org_vocab_size_padded
|
||||
padded_org_vocab_start_index, padded_org_vocab_end_index = (
|
||||
vocab_range_from_global_vocab_size(org_vocab_size_padded, tp_rank,
|
||||
tp_size))
|
||||
padded_added_vocab_start_index, padded_added_vocab_end_index = (
|
||||
vocab_range_from_global_vocab_size(num_added_embeddings_padded,
|
||||
tp_rank,
|
||||
tp_size,
|
||||
offset=org_vocab_size))
|
||||
# remove padding
|
||||
org_vocab_start_index = min(padded_org_vocab_start_index,
|
||||
org_vocab_size)
|
||||
org_vocab_end_index = min(padded_org_vocab_end_index, org_vocab_size)
|
||||
added_vocab_start_index = min(padded_added_vocab_start_index,
|
||||
vocab_size)
|
||||
added_vocab_end_index = min(padded_added_vocab_end_index, vocab_size)
|
||||
return VocabParallelEmbeddingShardIndices(
|
||||
padded_org_vocab_start_index, padded_org_vocab_end_index,
|
||||
padded_added_vocab_start_index, padded_added_vocab_end_index,
|
||||
org_vocab_start_index, org_vocab_end_index,
|
||||
added_vocab_start_index, added_vocab_end_index)
|
||||
|
||||
def get_sharded_to_full_mapping(self) -> Optional[List[int]]:
|
||||
"""Get a mapping that can be used to reindex the gathered
|
||||
logits for sampling.
|
||||
|
||||
During sampling, we gather logits from all ranks. The relationship
|
||||
of index->token_id will follow the same format as outlined in the class
|
||||
docstring. However, after the gather, we want to reindex the final
|
||||
logits tensor to map index->token_id one-to-one (the index is always
|
||||
equal the token_id it corresponds to). The indices returned by this
|
||||
method allow us to do that.
|
||||
"""
|
||||
if self.tp_size < 2:
|
||||
return None
|
||||
|
||||
base_embeddings: List[int] = []
|
||||
added_embeddings: List[int] = []
|
||||
padding: List[int] = []
|
||||
for tp_rank in range(self.tp_size):
|
||||
shard_indices = self._get_indices(self.num_embeddings_padded,
|
||||
self.org_vocab_size_padded,
|
||||
self.num_embeddings,
|
||||
self.org_vocab_size, tp_rank,
|
||||
self.tp_size)
|
||||
range_start = self.num_embeddings_per_partition * tp_rank
|
||||
range_end = self.num_embeddings_per_partition * (tp_rank + 1)
|
||||
base_embeddings.extend(
|
||||
range(range_start,
|
||||
range_start + shard_indices.num_org_elements))
|
||||
padding.extend(
|
||||
range(range_start + shard_indices.num_org_elements,
|
||||
range_start + shard_indices.num_org_elements_padded))
|
||||
added_embeddings.extend(
|
||||
range(
|
||||
range_start + shard_indices.num_org_elements_padded,
|
||||
range_start + shard_indices.num_org_elements_padded +
|
||||
shard_indices.num_added_elements))
|
||||
padding.extend(
|
||||
range(
|
||||
range_start + shard_indices.num_org_elements_padded +
|
||||
shard_indices.num_added_elements,
|
||||
range_start + shard_indices.num_org_elements_padded +
|
||||
shard_indices.num_added_elements_padded))
|
||||
assert (range_start + shard_indices.num_org_elements_padded +
|
||||
shard_indices.num_added_elements_padded == range_end)
|
||||
ret = base_embeddings + added_embeddings + padding
|
||||
assert len(ret) == self.num_embeddings_padded
|
||||
return ret
|
||||
|
||||
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
packed_dim = getattr(param, "packed_dim", None)
|
||||
|
||||
# If the parameter is a gguf weight, then load it directly.
|
||||
if getattr(param, "is_gguf_weight_type", None):
|
||||
param.data.copy_(loaded_weight)
|
||||
param.weight_type = loaded_weight.item()
|
||||
return
|
||||
elif isinstance(param, UninitializedParameter):
|
||||
shape = list(loaded_weight.shape)
|
||||
if output_dim is not None:
|
||||
shape[output_dim] = self.num_embeddings_per_partition
|
||||
param.materialize(tuple(shape), dtype=loaded_weight.dtype)
|
||||
|
||||
# If parameter does not have output dim, then it should
|
||||
# be copied onto all gpus (e.g. g_idx for act_order gptq).
|
||||
if output_dim is None:
|
||||
assert param.data.shape == loaded_weight.shape
|
||||
param.data.copy_(loaded_weight)
|
||||
return
|
||||
|
||||
# Shard indexes for loading the weight
|
||||
start_idx = self.shard_indices.org_vocab_start_index
|
||||
shard_size = self.shard_indices.org_vocab_end_index - start_idx
|
||||
|
||||
# If param packed on the same dim we are sharding on, then
|
||||
# need to adjust offsets of loaded weight by pack_factor.
|
||||
if packed_dim is not None and packed_dim == output_dim:
|
||||
packed_factor = param.packed_factor if isinstance(
|
||||
param, BasevLLMParameter) else param.pack_factor
|
||||
assert loaded_weight.shape[output_dim] == (self.org_vocab_size //
|
||||
param.packed_factor)
|
||||
start_idx = start_idx // packed_factor
|
||||
shard_size = shard_size // packed_factor
|
||||
else:
|
||||
assert loaded_weight.shape[output_dim] == self.org_vocab_size
|
||||
|
||||
# Copy the data. Select chunk corresponding to current shard.
|
||||
loaded_weight = loaded_weight.narrow(output_dim, start_idx, shard_size)
|
||||
|
||||
if current_platform.is_hpu():
|
||||
# FIXME(kzawora): Weight copy with slicing bugs out on Gaudi here,
|
||||
# so we're using a workaround. Remove this when fixed in
|
||||
# HPU PT bridge.
|
||||
padded_weight = torch.cat([
|
||||
loaded_weight,
|
||||
torch.zeros(param.shape[0] - loaded_weight.shape[0],
|
||||
*loaded_weight.shape[1:])
|
||||
])
|
||||
param.data.copy_(padded_weight)
|
||||
else:
|
||||
param[:loaded_weight.shape[0]].data.copy_(loaded_weight)
|
||||
param[loaded_weight.shape[0]:].data.fill_(0)
|
||||
|
||||
def forward(self, input_):
|
||||
if self.tp_size > 1:
|
||||
# Build the mask.
|
||||
masked_input, input_mask = get_masked_input_and_mask(
|
||||
input_, self.shard_indices.org_vocab_start_index,
|
||||
self.shard_indices.org_vocab_end_index,
|
||||
self.shard_indices.num_org_vocab_padding,
|
||||
self.shard_indices.added_vocab_start_index,
|
||||
self.shard_indices.added_vocab_end_index)
|
||||
else:
|
||||
masked_input = input_
|
||||
# Get the embeddings.
|
||||
output_parallel = self.quant_method.embedding(self,
|
||||
masked_input.long())
|
||||
# Mask the output embedding.
|
||||
if self.tp_size > 1:
|
||||
output_parallel.masked_fill_(input_mask.unsqueeze(-1), 0)
|
||||
# Reduce across all the model parallel GPUs.
|
||||
output = tensor_model_parallel_all_reduce(output_parallel)
|
||||
return output
|
||||
|
||||
def extra_repr(self) -> str:
|
||||
s = f"num_embeddings={self.num_embeddings_per_partition}"
|
||||
s += f", embedding_dim={self.embedding_dim}"
|
||||
s += f", org_vocab_size={self.org_vocab_size}"
|
||||
s += f', num_embeddings_padded={self.num_embeddings_padded}'
|
||||
s += f', tp_size={self.tp_size}'
|
||||
return s
|
||||
|
||||
|
||||
class ParallelLMHead(VocabParallelEmbedding):
|
||||
"""Parallelized LM head.
|
||||
|
||||
Output logits weight matrices used in the Sampler. The weight and bias
|
||||
tensors are padded to make sure they are divisible by the number of
|
||||
model parallel GPUs.
|
||||
|
||||
Args:
|
||||
num_embeddings: vocabulary size.
|
||||
embedding_dim: size of hidden state.
|
||||
bias: whether to use bias.
|
||||
params_dtype: type of the parameters.
|
||||
org_num_embeddings: original vocabulary size (without LoRA).
|
||||
padding_size: padding size for the vocabulary.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
num_embeddings: int,
|
||||
embedding_dim: int,
|
||||
bias: bool = False,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
org_num_embeddings: Optional[int] = None,
|
||||
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__(num_embeddings, embedding_dim, params_dtype,
|
||||
org_num_embeddings, padding_size, quant_config,
|
||||
prefix)
|
||||
self.quant_config = quant_config
|
||||
if bias:
|
||||
self.bias = Parameter(
|
||||
torch.empty(self.num_embeddings_per_partition,
|
||||
dtype=params_dtype))
|
||||
set_weight_attrs(self.bias, {
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
})
|
||||
else:
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
def tie_weights(self, embed_tokens: VocabParallelEmbedding):
|
||||
"""Tie the weights with word embeddings."""
|
||||
# GGUF quantized embed_tokens.
|
||||
if self.quant_config and self.quant_config.get_name() == "gguf":
|
||||
return embed_tokens
|
||||
else:
|
||||
self.weight = embed_tokens.weight
|
||||
return self
|
||||
|
||||
def forward(self, input_):
|
||||
del input_
|
||||
raise RuntimeError("LMHead's weights should be used in the sampler.")
|
||||
@@ -0,0 +1,219 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm
|
||||
# https://github.com/vllm-project/vllm/blob/main/vllm/logger.py
|
||||
# Copyright 2023 The vLLM Authors.
|
||||
# Copyright 2023 The FastVideo Authors.
|
||||
|
||||
"""Logging configuration for fastvideo.v1."""
|
||||
import datetime
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from functools import lru_cache, partial
|
||||
from logging import Logger
|
||||
from logging.config import dictConfig
|
||||
from os import path
|
||||
from types import MethodType
|
||||
from typing import Any, Optional, cast
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
|
||||
FASTVIDEO_CONFIGURE_LOGGING = envs.FASTVIDEO_CONFIGURE_LOGGING
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH = envs.FASTVIDEO_LOGGING_CONFIG_PATH
|
||||
FASTVIDEO_LOGGING_LEVEL = envs.FASTVIDEO_LOGGING_LEVEL
|
||||
FASTVIDEO_LOGGING_PREFIX = envs.FASTVIDEO_LOGGING_PREFIX
|
||||
|
||||
_FORMAT = (f"{FASTVIDEO_LOGGING_PREFIX}%(levelname)s %(asctime)s "
|
||||
"[%(filename)s:%(lineno)d] %(message)s")
|
||||
_DATE_FORMAT = "%m-%d %H:%M:%S"
|
||||
|
||||
DEFAULT_LOGGING_CONFIG = {
|
||||
"formatters": {
|
||||
"fastvideo": {
|
||||
"class": "fastvideo.v1.logging_utils.NewLineFormatter",
|
||||
"datefmt": _DATE_FORMAT,
|
||||
"format": _FORMAT,
|
||||
},
|
||||
},
|
||||
"handlers": {
|
||||
"fastvideo": {
|
||||
"class": "logging.StreamHandler",
|
||||
"formatter": "fastvideo",
|
||||
"level": FASTVIDEO_LOGGING_LEVEL,
|
||||
"stream": "ext://sys.stdout",
|
||||
},
|
||||
},
|
||||
"loggers": {
|
||||
"fastvideo": {
|
||||
"handlers": ["fastvideo"],
|
||||
"level": "DEBUG",
|
||||
"propagate": False,
|
||||
},
|
||||
},
|
||||
"root": {
|
||||
"handlers": ["fastvideo"],
|
||||
"level": "DEBUG",
|
||||
},
|
||||
"version": 1,
|
||||
"disable_existing_loggers": False
|
||||
}
|
||||
|
||||
|
||||
@lru_cache
|
||||
def _print_info_once(logger: Logger, msg: str) -> None:
|
||||
# Set the stacklevel to 2 to print the original caller's line info
|
||||
logger.info(msg, stacklevel=2)
|
||||
|
||||
|
||||
@lru_cache
|
||||
def _print_warning_once(logger: Logger, msg: str) -> None:
|
||||
# Set the stacklevel to 2 to print the original caller's line info
|
||||
logger.warning(msg, stacklevel=2)
|
||||
|
||||
|
||||
class _FastvideoLogger(Logger):
|
||||
"""
|
||||
Note:
|
||||
This class is just to provide type information.
|
||||
We actually patch the methods directly on the :class:`logging.Logger`
|
||||
instance to avoid conflicting with other libraries such as
|
||||
`intel_extension_for_pytorch.utils._logger`.
|
||||
"""
|
||||
|
||||
def info_once(self, msg: str) -> None:
|
||||
"""
|
||||
As :meth:`info`, but subsequent calls with the same message
|
||||
are silently dropped.
|
||||
"""
|
||||
_print_info_once(self, msg)
|
||||
|
||||
def warning_once(self, msg: str) -> None:
|
||||
"""
|
||||
As :meth:`warning`, but subsequent calls with the same message
|
||||
are silently dropped.
|
||||
"""
|
||||
_print_warning_once(self, msg)
|
||||
|
||||
|
||||
def _configure_fastvideo_root_logger() -> None:
|
||||
logging_config = dict[str, Any]()
|
||||
|
||||
if not FASTVIDEO_CONFIGURE_LOGGING and FASTVIDEO_LOGGING_CONFIG_PATH:
|
||||
raise RuntimeError(
|
||||
"FASTVIDEO_CONFIGURE_LOGGING evaluated to false, but "
|
||||
"FASTVIDEO_LOGGING_CONFIG_PATH was given. FASTVIDEO_LOGGING_CONFIG_PATH "
|
||||
"implies FASTVIDEO_CONFIGURE_LOGGING. Please enable "
|
||||
"FASTVIDEO_CONFIGURE_LOGGING or unset FASTVIDEO_LOGGING_CONFIG_PATH.")
|
||||
|
||||
if FASTVIDEO_CONFIGURE_LOGGING:
|
||||
logging_config = DEFAULT_LOGGING_CONFIG
|
||||
|
||||
if FASTVIDEO_LOGGING_CONFIG_PATH:
|
||||
if not path.exists(FASTVIDEO_LOGGING_CONFIG_PATH):
|
||||
raise RuntimeError(
|
||||
"Could not load logging config. File does not exist: %s",
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH)
|
||||
with open(FASTVIDEO_LOGGING_CONFIG_PATH, encoding="utf-8") as file:
|
||||
custom_config = json.loads(file.read())
|
||||
|
||||
if not isinstance(custom_config, dict):
|
||||
raise ValueError("Invalid logging config. Expected Dict, got %s.",
|
||||
type(custom_config).__name__)
|
||||
logging_config = custom_config
|
||||
|
||||
for formatter in logging_config.get("formatters", {}).values():
|
||||
# This provides backwards compatibility after #10134.
|
||||
if formatter.get("class") == "fastvideo.v1.logging.NewLineFormatter":
|
||||
formatter["class"] = "fastvideo.v1.logging_utils.NewLineFormatter"
|
||||
|
||||
if logging_config:
|
||||
dictConfig(logging_config)
|
||||
|
||||
# TODO: add rank_zero_only log
|
||||
def init_logger(name: str) -> _FastvideoLogger:
|
||||
"""The main purpose of this function is to ensure that loggers are
|
||||
retrieved in such a way that we can be sure the root fastvideo logger has
|
||||
already been configured."""
|
||||
|
||||
logger = logging.getLogger(name)
|
||||
|
||||
methods_to_patch = {
|
||||
"info_once": _print_info_once,
|
||||
"warning_once": _print_warning_once,
|
||||
}
|
||||
|
||||
for method_name, method in methods_to_patch.items():
|
||||
setattr(logger, method_name, MethodType(method, logger))
|
||||
|
||||
return cast(_FastvideoLogger, logger)
|
||||
|
||||
|
||||
# The root logger is initialized when the module is imported.
|
||||
# This is thread-safe as the module is only imported once,
|
||||
# guaranteed by the Python GIL.
|
||||
_configure_fastvideo_root_logger()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _trace_calls(log_path, root_dir, frame, event, arg=None):
|
||||
if event in ['call', 'return']:
|
||||
# Extract the filename, line number, function name, and the code object
|
||||
filename = frame.f_code.co_filename
|
||||
lineno = frame.f_lineno
|
||||
func_name = frame.f_code.co_name
|
||||
if not filename.startswith(root_dir):
|
||||
# only log the functions in the fastvideo root_dir
|
||||
return
|
||||
# Log every function call or return
|
||||
try:
|
||||
last_frame = frame.f_back
|
||||
if last_frame is not None:
|
||||
last_filename = last_frame.f_code.co_filename
|
||||
last_lineno = last_frame.f_lineno
|
||||
last_func_name = last_frame.f_code.co_name
|
||||
else:
|
||||
# initial frame
|
||||
last_filename = ""
|
||||
last_lineno = 0
|
||||
last_func_name = ""
|
||||
with open(log_path, 'a') as f:
|
||||
ts = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S.%f")
|
||||
if event == 'call':
|
||||
f.write(f"{ts} Call to"
|
||||
f" {func_name} in {filename}:{lineno}"
|
||||
f" from {last_func_name} in {last_filename}:"
|
||||
f"{last_lineno}\n")
|
||||
else:
|
||||
f.write(f"{ts} Return from"
|
||||
f" {func_name} in {filename}:{lineno}"
|
||||
f" to {last_func_name} in {last_filename}:"
|
||||
f"{last_lineno}\n")
|
||||
except NameError:
|
||||
# modules are deleted during shutdown
|
||||
pass
|
||||
return partial(_trace_calls, log_path, root_dir)
|
||||
|
||||
|
||||
def enable_trace_function_call(log_file_path: str,
|
||||
root_dir: Optional[str] = None):
|
||||
"""
|
||||
Enable tracing of every function call in code under `root_dir`.
|
||||
This is useful for debugging hangs or crashes.
|
||||
`log_file_path` is the path to the log file.
|
||||
`root_dir` is the root directory of the code to trace. If None, it is the
|
||||
fastvideo root directory.
|
||||
|
||||
Note that this call is thread-level, any threads calling this function
|
||||
will have the trace enabled. Other threads will not be affected.
|
||||
"""
|
||||
logger.warning(
|
||||
"FASTVIDEO_TRACE_FUNCTION is enabled. It will record every"
|
||||
" function executed by Python. This will slow down the code. It "
|
||||
"is suggested to be used for debugging hang or crashes only.")
|
||||
logger.info("Trace frame log is saved to %s", log_file_path)
|
||||
if root_dir is None:
|
||||
# by default, this is the fastvideo root directory
|
||||
root_dir = os.path.dirname(os.path.dirname(__file__))
|
||||
sys.settrace(partial(_trace_calls, log_file_path, root_dir))
|
||||
@@ -0,0 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
|
||||
from fastvideo.v1.logging_utils.formatter import NewLineFormatter
|
||||
|
||||
__all__ = [
|
||||
"NewLineFormatter",
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm
|
||||
|
||||
import logging
|
||||
|
||||
|
||||
class NewLineFormatter(logging.Formatter):
|
||||
"""Adds logging prefix to newlines to align multi-line messages."""
|
||||
|
||||
def __init__(self, fmt, datefmt=None, style="%"):
|
||||
logging.Formatter.__init__(self, fmt, datefmt, style)
|
||||
|
||||
def format(self, record):
|
||||
msg = logging.Formatter.format(self, record)
|
||||
if record.message != "":
|
||||
parts = msg.split(record.message)
|
||||
msg = msg.replace("\n", "\r\n" + parts[0])
|
||||
return msg
|
||||
@@ -0,0 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch.nn as nn
|
||||
from typing import Dict
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def get_scheduler(module_path: str, architecture: str, inference_args: InferenceArgs) -> Dict:
|
||||
"""Create a scheduler based on the inference args. Can be overridden by subclasses."""
|
||||
if hasattr(inference_args, 'denoise_type') and inference_args.denoise_type == "flow":
|
||||
# TODO(will): add schedulers to register or create a new scheduler registry
|
||||
# TODO(will): default to config file but allow override through
|
||||
# inference args. Currently only uses inference args.
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchDiscreteScheduler
|
||||
return FlowMatchDiscreteScheduler(
|
||||
shift=inference_args.flow_shift,
|
||||
solver=inference_args.flow_solver,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Invalid denoise type: {inference_args.denoise_type}")
|
||||
|
||||
__all__ = [
|
||||
"set_random_seed",
|
||||
"BasevLLMParameter",
|
||||
"PackedvLLMParameter",
|
||||
"get_model",
|
||||
"get_scheduler",
|
||||
]
|
||||
@@ -0,0 +1,12 @@
|
||||
from torch import nn
|
||||
|
||||
|
||||
class BaseDiT(nn.Module):
|
||||
_fsdp_shard_conditions = []
|
||||
attention_head_dim: int = None
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
pass
|
||||
@@ -0,0 +1,897 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from typing import Optional, Tuple, List
|
||||
from fastvideo.v1.attention.flash_attn import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
from fastvideo.v1.layers.layernorm import LayerNormScaleShift, ScaleResidual, ScaleResidualLayerNormScaleShift
|
||||
from fastvideo.v1.layers.visual_embedding import PatchEmbed, TimestepEmbedder, ModulateProjection, unpatchify
|
||||
from fastvideo.v1.layers.rotary_embedding import _apply_rotary_emb, get_rotary_pos_embed
|
||||
from fastvideo.v1.distributed.parallel_state import get_sequence_model_parallel_world_size, get_sequence_model_parallel_rank
|
||||
# TODO(will-PY-refactor): RMSNorm ....
|
||||
from fastvideo.v1.layers.mlp import MLP
|
||||
from fastvideo.v1.models.dits.base import BaseDiT
|
||||
from diffusers.loaders import PeftAdapterMixin
|
||||
|
||||
class HunyuanRMSNorm(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
elementwise_affine=True,
|
||||
eps: float = 1e-6,
|
||||
device=None,
|
||||
dtype=None,
|
||||
):
|
||||
"""
|
||||
Initialize the RMSNorm normalization layer.
|
||||
|
||||
Args:
|
||||
dim (int): The dimension of the input tensor.
|
||||
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
|
||||
|
||||
Attributes:
|
||||
eps (float): A small value added to the denominator for numerical stability.
|
||||
weight (nn.Parameter): Learnable scaling parameter.
|
||||
|
||||
"""
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
if elementwise_affine:
|
||||
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
|
||||
|
||||
def _norm(self, x):
|
||||
"""
|
||||
Apply the RMSNorm normalization to the input tensor.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): The input tensor.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The normalized tensor.
|
||||
|
||||
"""
|
||||
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
|
||||
|
||||
def forward(self, x):
|
||||
"""
|
||||
Forward pass through the RMSNorm layer.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): The input tensor.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The output tensor after applying RMSNorm.
|
||||
|
||||
"""
|
||||
output = self._norm(x.float()).type_as(x)
|
||||
if hasattr(self, "weight"):
|
||||
output = output * self.weight
|
||||
return output
|
||||
|
||||
class MMDoubleStreamBlock(nn.Module):
|
||||
"""
|
||||
A multimodal DiT block with separate modulation for text and image/video,
|
||||
using distributed attention and linear layers.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.deterministic = False
|
||||
self.num_attention_heads = num_attention_heads
|
||||
head_dim = hidden_size // num_attention_heads
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
|
||||
# Image modulation components
|
||||
self.img_mod = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=6,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
# Fused operations for image stream
|
||||
self.img_attn_norm = LayerNormScaleShift(hidden_size, norm_type="layer", elementwise_affine=False, dtype=dtype)
|
||||
self.img_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(hidden_size, norm_type="layer", elementwise_affine=False, dtype=dtype)
|
||||
self.img_mlp_residual = ScaleResidual()
|
||||
|
||||
# Image attention components
|
||||
# self.img_attn_qkv = ReplicatedLinear(
|
||||
# hidden_size,
|
||||
# hidden_size * 3,
|
||||
# bias=True,
|
||||
# params_dtype=dtype
|
||||
# )
|
||||
for name in ["q", "k", "v"]:
|
||||
setattr(self, f"img_attn_to_{name}", ReplicatedLinear(
|
||||
hidden_size, hidden_size, bias=True, params_dtype=dtype
|
||||
))
|
||||
|
||||
self.img_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
self.img_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
|
||||
|
||||
self.img_attn_proj = ReplicatedLinear(
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
self.img_mlp = MLP(
|
||||
hidden_size,
|
||||
mlp_hidden_dim,
|
||||
bias=True,
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
# Text modulation components
|
||||
self.txt_mod = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=6,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
# Fused operations for text stream
|
||||
self.txt_attn_norm = LayerNormScaleShift(hidden_size, norm_type="layer", elementwise_affine=False, dtype=dtype)
|
||||
self.txt_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(hidden_size, norm_type="layer", elementwise_affine=False, dtype=dtype)
|
||||
self.txt_mlp_residual = ScaleResidual()
|
||||
|
||||
# Text attention components
|
||||
self.txt_attn_qkv = ReplicatedLinear(
|
||||
hidden_size,
|
||||
hidden_size * 3,
|
||||
bias=True,
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
# QK norm layers for text
|
||||
self.txt_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
self.txt_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
|
||||
self.txt_attn_proj = ReplicatedLinear(
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
self.txt_mlp = MLP(
|
||||
hidden_size,
|
||||
mlp_hidden_dim,
|
||||
bias=True,
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
# Distributed attention
|
||||
self.attn = DistributedAttention(
|
||||
dropout_rate=0.0,
|
||||
causal=False
|
||||
)
|
||||
|
||||
|
||||
def forward(
|
||||
self,
|
||||
img: torch.Tensor,
|
||||
txt: torch.Tensor,
|
||||
vec: torch.Tensor,
|
||||
freqs_cis: tuple = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# Process modulation vectors
|
||||
img_mod_outputs = self.img_mod(vec)
|
||||
(
|
||||
img_attn_shift,
|
||||
img_attn_scale,
|
||||
img_attn_gate,
|
||||
img_mlp_shift,
|
||||
img_mlp_scale,
|
||||
img_mlp_gate,
|
||||
) = torch.chunk(img_mod_outputs, 6, dim=-1)
|
||||
|
||||
txt_mod_outputs = self.txt_mod(vec)
|
||||
(
|
||||
txt_attn_shift,
|
||||
txt_attn_scale,
|
||||
txt_attn_gate,
|
||||
txt_mlp_shift,
|
||||
txt_mlp_scale,
|
||||
txt_mlp_gate,
|
||||
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
|
||||
|
||||
# Prepare image for attention using fused operation
|
||||
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale)
|
||||
# Get QKV for image
|
||||
# img_qkv, _ = self.img_attn_qkv(img_attn_input)
|
||||
img_q_, _ = self.img_attn_to_q(img_attn_input)
|
||||
img_k_, _ = self.img_attn_to_k(img_attn_input)
|
||||
img_v_, _ = self.img_attn_to_v(img_attn_input)
|
||||
img_qkv = torch.cat((img_q_, img_k_, img_v_), dim=-1)
|
||||
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
|
||||
# batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
|
||||
|
||||
# Split QKV
|
||||
img_qkv = img_qkv.view(batch_size, image_seq_len, 3, self.num_attention_heads, -1)
|
||||
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :, 2]
|
||||
|
||||
# Apply QK-Norm if needed
|
||||
|
||||
img_q = self.img_attn_q_norm(img_q).to(img_v)
|
||||
img_k = self.img_attn_k_norm(img_k).to(img_v)
|
||||
# Apply rotary embeddings
|
||||
cos, sin = freqs_cis
|
||||
img_q, img_k = _apply_rotary_emb(img_q, cos, sin, is_neox_style=False), _apply_rotary_emb(img_k, cos, sin, is_neox_style=False)
|
||||
# Prepare text for attention using fused operation
|
||||
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
|
||||
|
||||
# Get QKV for text
|
||||
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
|
||||
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
|
||||
|
||||
# Split QKV
|
||||
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3, self.num_attention_heads, -1)
|
||||
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :, 2]
|
||||
|
||||
# Apply QK-Norm if needed
|
||||
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
|
||||
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
|
||||
|
||||
# Run distributed attention
|
||||
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v)
|
||||
img_attn_out, _ = self.img_attn_proj(img_attn.view(batch_size, image_seq_len, -1))
|
||||
# Use fused operation for residual connection, normalization, and modulation
|
||||
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
|
||||
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale
|
||||
)
|
||||
|
||||
# Process image MLP
|
||||
img_mlp_out = self.img_mlp(img_mlp_input)
|
||||
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
|
||||
|
||||
# Process text attention output
|
||||
txt_attn_out, _ = self.txt_attn_proj(txt_attn.reshape(batch_size, text_seq_len, -1))
|
||||
|
||||
# Use fused operation for residual connection, normalization, and modulation
|
||||
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
|
||||
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale
|
||||
)
|
||||
|
||||
# Process text MLP
|
||||
txt_mlp_out = self.txt_mlp(txt_mlp_input)
|
||||
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
|
||||
|
||||
return img, txt
|
||||
|
||||
|
||||
class MMSingleStreamBlock(nn.Module):
|
||||
"""
|
||||
A DiT block with parallel linear layers using distributed attention
|
||||
and tensor parallelism.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.deterministic = False
|
||||
self.hidden_size = hidden_size
|
||||
self.num_attention_heads = num_attention_heads
|
||||
head_dim = hidden_size // num_attention_heads
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
self.mlp_hidden_dim = mlp_hidden_dim
|
||||
|
||||
# Combined QKV and MLP input projection
|
||||
# self.linear1 = ReplicatedLinear(
|
||||
# hidden_size,
|
||||
# hidden_size * 3 + mlp_hidden_dim,
|
||||
# bias=True,
|
||||
# params_dtype=dtype
|
||||
# )
|
||||
# self.linear1_qkv = ReplicatedLinear(
|
||||
# hidden_size,
|
||||
# hidden_size * 3,
|
||||
# bias=True,
|
||||
# params_dtype=dtype
|
||||
# )
|
||||
for name in ["q", "k", "v"]:
|
||||
setattr(self, f"linear1_to_{name}", ReplicatedLinear(
|
||||
hidden_size, hidden_size, bias=True, params_dtype=dtype
|
||||
))
|
||||
self.linear1_mlp = ReplicatedLinear(
|
||||
hidden_size,
|
||||
mlp_hidden_dim,
|
||||
bias=True,
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
# Combined projection and MLP output
|
||||
self.linear2 = ReplicatedLinear(
|
||||
hidden_size + mlp_hidden_dim,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
# QK norm layers
|
||||
self.q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
self.k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
|
||||
|
||||
# Fused operations with better naming
|
||||
self.input_norm_scale_shift = LayerNormScaleShift(hidden_size, norm_type="layer", eps=1e-6, elementwise_affine=False, dtype=dtype)
|
||||
self.output_residual = ScaleResidual()
|
||||
|
||||
# Activation function
|
||||
self.mlp_act = nn.GELU(approximate="tanh")
|
||||
|
||||
|
||||
# Modulation
|
||||
self.modulation = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=3,
|
||||
act_layer="silu",
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
# Distributed attention
|
||||
self.attn = DistributedAttention(
|
||||
dropout_rate=0.0,
|
||||
causal=False
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
vec: torch.Tensor,
|
||||
txt_len: int,
|
||||
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
# Process modulation
|
||||
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
|
||||
|
||||
# Apply pre-norm and modulation using fused operation
|
||||
x_mod = self.input_norm_scale_shift(x, mod_shift, mod_scale)
|
||||
|
||||
# Get combined projections
|
||||
linear1_q, _ = self.linear1_to_q(x_mod)
|
||||
linear1_k, _ = self.linear1_to_k(x_mod)
|
||||
linear1_v, _ = self.linear1_to_v(x_mod)
|
||||
linear1_qkv = torch.cat((linear1_q, linear1_k, linear1_v), dim=-1)
|
||||
linear1_mlp, _ = self.linear1_mlp(x_mod)
|
||||
linear1_out = torch.cat((linear1_qkv, linear1_mlp), dim=-1)
|
||||
|
||||
# Split into QKV and MLP parts
|
||||
qkv, mlp = torch.split(
|
||||
linear1_out, [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
|
||||
)
|
||||
|
||||
# Process QKV
|
||||
batch_size, seq_len = qkv.shape[0], qkv.shape[1]
|
||||
qkv = qkv.view(batch_size, seq_len, 3, self.num_attention_heads, -1)
|
||||
q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2]
|
||||
|
||||
# Apply QK-Norm
|
||||
q = self.q_norm(q).to(v.dtype)
|
||||
k = self.k_norm(k).to(v.dtype)
|
||||
|
||||
|
||||
# Split into image and text parts
|
||||
img_q, txt_q = q[:, :-txt_len], q[:, -txt_len:]
|
||||
img_k, txt_k = k[:, :-txt_len], k[:, -txt_len:]
|
||||
img_v, txt_v = v[:, :-txt_len], v[:, -txt_len:]
|
||||
# Apply rotary embeddings to image parts
|
||||
cos, sin = freqs_cis
|
||||
img_q, img_k = _apply_rotary_emb(img_q, cos, sin, is_neox_style=False), _apply_rotary_emb(img_k, cos, sin, is_neox_style=False)
|
||||
|
||||
|
||||
# Run distributed attention
|
||||
img_attn_output, txt_attn_output = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v)
|
||||
attn_output = torch.cat((img_attn_output, txt_attn_output), dim=1).view(batch_size, seq_len, -1)
|
||||
# Process MLP activation
|
||||
mlp_output = self.mlp_act(mlp)
|
||||
|
||||
# Combine attention and MLP outputs
|
||||
combined = torch.cat((attn_output, mlp_output), dim=-1)
|
||||
|
||||
# Final projection
|
||||
output, _ = self.linear2(combined)
|
||||
|
||||
# Apply residual connection with gating using fused operation
|
||||
return self.output_residual(x, output, mod_gate)
|
||||
|
||||
|
||||
|
||||
|
||||
class HunyuanVideoTransformer3DModel(BaseDiT, PeftAdapterMixin):
|
||||
"""
|
||||
HunyuanVideo Transformer backbone adapted for distributed training.
|
||||
|
||||
This implementation uses distributed attention and linear layers for efficient
|
||||
parallel processing across multiple GPUs.
|
||||
|
||||
Based on the architecture from:
|
||||
- Flux.1: https://github.com/black-forest-labs/flux
|
||||
- MMDiT: http://arxiv.org/abs/2403.03206
|
||||
"""
|
||||
# PY: we make the input args the same as HF config
|
||||
|
||||
# shard single stream, double stream blocks, and refiner_blocks
|
||||
_fsdp_shard_conditions = [
|
||||
lambda n, m: "double" in n and str.isdigit(n.split(".")[-1]),
|
||||
lambda n, m: "single" in n and str.isdigit(n.split(".")[-1]),
|
||||
lambda n, m: "refiner" in n and str.isdigit(n.split(".")[-1]),
|
||||
]
|
||||
_param_names_mapping = {
|
||||
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
|
||||
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$": r"txt_in.t_embedder.mlp.fc_in.\1",
|
||||
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$": r"txt_in.t_embedder.mlp.fc_out.\1",
|
||||
r"^context_embedder\.proj_in\.(.*)$": r"txt_in.input_embedder.\1",
|
||||
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$": r"txt_in.c_embedder.fc_in.\1",
|
||||
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$": r"txt_in.c_embedder.fc_out.\1",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$": r"txt_in.refiner_blocks.\1.norm1.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$": r"txt_in.refiner_blocks.\1.norm2.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$": r"txt_in.refiner_blocks.\1.self_attn_to_q.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$": r"txt_in.refiner_blocks.\1.self_attn_to_k.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$": r"txt_in.refiner_blocks.\1.self_attn_to_v.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$": r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$": r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$": r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$": r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
|
||||
|
||||
# 3. x_embedder mapping:
|
||||
r"^x_embedder\.proj\.(.*)$": r"img_in.proj.\1",
|
||||
|
||||
# 4. Top-level time_text_embed mappings:
|
||||
r"^time_text_embed\.timestep_embedder\.linear_1\.(.*)$": r"time_in.mlp.fc_in.\1",
|
||||
r"^time_text_embed\.timestep_embedder\.linear_2\.(.*)$": r"time_in.mlp.fc_out.\1",
|
||||
r"^time_text_embed\.guidance_embedder\.linear_1\.(.*)$": r"guidance_in.mlp.fc_in.\1",
|
||||
r"^time_text_embed\.guidance_embedder\.linear_2\.(.*)$": r"guidance_in.mlp.fc_out.\1",
|
||||
r"^time_text_embed\.text_embedder\.linear_1\.(.*)$": r"vector_in.fc_in.\1",
|
||||
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$": r"vector_in.fc_out.\1",
|
||||
|
||||
# 5. transformer_blocks mapping:
|
||||
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$": r"double_blocks.\1.img_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$": r"double_blocks.\1.txt_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$": r"double_blocks.\1.img_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$": r"double_blocks.\1.img_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": r"double_blocks.\1.img_attn_to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$": r"double_blocks.\1.img_attn_to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$": r"double_blocks.\1.img_attn_to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$": r"double_blocks.\1.img_attn_proj.\2",
|
||||
# Corrected: merge attn.to_add_out into the main projection.
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$": r"double_blocks.\1.txt_attn_proj.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$": r"double_blocks.\1.txt_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$": r"double_blocks.\1.txt_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$": r"double_blocks.\1.img_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$": r"double_blocks.\1.img_mlp.fc_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$": r"double_blocks.\1.txt_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$": r"double_blocks.\1.txt_mlp.fc_out.\2",
|
||||
|
||||
# 6. single_transformer_blocks mapping:
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$": r"single_blocks.\1.q_norm.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$": r"single_blocks.\1.k_norm.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": r"single_blocks.\1.linear1_to_q.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$": r"single_blocks.\1.linear1_to_k.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$": r"single_blocks.\1.linear1_to_v.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$": r"single_blocks.\1.linear1_mlp.\2",
|
||||
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
|
||||
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$": r"single_blocks.\1.linear2.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.norm\.linear\.(.*)$": r"single_blocks.\1.modulation.linear.\2",
|
||||
|
||||
# 7. Final layers mapping:
|
||||
r"^norm_out\.linear\.(.*)$": r"final_layer.adaLN_modulation.linear.\1",
|
||||
r"^proj_out\.(.*)$": r"final_layer.linear.\1",
|
||||
}
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int = 2,
|
||||
patch_size_t: int = 1,
|
||||
in_channels: int = 16,
|
||||
out_channels: int = 16,
|
||||
num_attention_heads: int = 24,
|
||||
attention_head_dim: int = 128,
|
||||
mlp_ratio: float = 4.0,
|
||||
num_layers: int = 20,
|
||||
num_single_layers: int = 40,
|
||||
num_refiner_layers: int = 2,
|
||||
rope_axes_dim: List[int] = [16, 56, 56],
|
||||
guidance_embeds: bool = False,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
text_embed_dim: int = 4096,
|
||||
pooled_projection_dim: int = 768,
|
||||
rope_theta: int = 256,
|
||||
qk_norm: str = "rms_norm", #TODO(PY)
|
||||
):
|
||||
super().__init__()
|
||||
hidden_size = attention_head_dim * num_attention_heads
|
||||
self.patch_size = [patch_size_t, patch_size, patch_size]
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = in_channels if out_channels is None else out_channels
|
||||
self.unpatchify_channels = self.out_channels
|
||||
self.guidance_embeds = guidance_embeds
|
||||
self.rope_dim_list = rope_axes_dim
|
||||
self.rope_theta = rope_theta
|
||||
self.text_states_dim = text_embed_dim
|
||||
self.text_states_dim_2 = pooled_projection_dim
|
||||
|
||||
if hidden_size % num_attention_heads != 0:
|
||||
raise ValueError(f"Hidden size {hidden_size} must be divisible by num_attention_heads {num_attention_heads}")
|
||||
|
||||
pe_dim = hidden_size // num_attention_heads
|
||||
if sum(rope_axes_dim) != pe_dim:
|
||||
raise ValueError(f"Got {rope_axes_dim} but expected positional dim {pe_dim}")
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.num_attention_heads = num_attention_heads
|
||||
|
||||
# Image projection
|
||||
self.img_in = PatchEmbed(
|
||||
self.patch_size,
|
||||
self.in_channels,
|
||||
self.hidden_size,
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
|
||||
self.txt_in = SingleTokenRefiner(
|
||||
self.text_states_dim,
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
depth=num_refiner_layers,
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
|
||||
# Time modulation
|
||||
self.time_in = TimestepEmbedder(
|
||||
self.hidden_size,
|
||||
act_layer="silu",
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
# Text modulation
|
||||
self.vector_in = MLP(
|
||||
self.text_states_dim_2,
|
||||
self.hidden_size,
|
||||
self.hidden_size,
|
||||
act_type="silu",
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
# Guidance modulation
|
||||
self.guidance_in = (
|
||||
TimestepEmbedder(self.hidden_size, act_layer="silu", dtype=dtype)
|
||||
if guidance_embeds else None
|
||||
)
|
||||
|
||||
# Double blocks
|
||||
self.double_blocks = nn.ModuleList([
|
||||
MMDoubleStreamBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
dtype=dtype,
|
||||
) for _ in range(num_layers)
|
||||
])
|
||||
|
||||
# Single blocks
|
||||
self.single_blocks = nn.ModuleList([
|
||||
MMSingleStreamBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
dtype=dtype,
|
||||
) for _ in range(num_single_layers)
|
||||
])
|
||||
|
||||
self.final_layer = FinalLayer(
|
||||
hidden_size,
|
||||
self.patch_size,
|
||||
self.out_channels,
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
# TODO: change the input the FORWAD_BACTCH Dict
|
||||
# TODO: change output to a dict
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
guidance=None,
|
||||
):
|
||||
"""
|
||||
Forward pass of the HunyuanDiT model.
|
||||
|
||||
Args:
|
||||
hidden_states: Input image/video latents [B, C, T, H, W]
|
||||
encoder_hidden_states: Text embeddings [B, L, D]
|
||||
timestep: Diffusion timestep
|
||||
guidance: Guidance scale for CFG
|
||||
|
||||
Returns:
|
||||
Tuple of (output)
|
||||
"""
|
||||
if guidance is None:
|
||||
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=hidden_states.dtype)
|
||||
|
||||
img = x = hidden_states
|
||||
t = timestep
|
||||
|
||||
# Split text embeddings - first token is global, rest are per-token
|
||||
txt = encoder_hidden_states[:, 1:]
|
||||
text_states_2 = encoder_hidden_states[:, 0, :self.text_states_dim_2]
|
||||
|
||||
# Get spatial dimensions
|
||||
_, _, ot, oh, ow = x.shape
|
||||
tt, th, tw = (
|
||||
ot // self.patch_size[0],
|
||||
oh // self.patch_size[1],
|
||||
ow // self.patch_size[2],
|
||||
)
|
||||
|
||||
# Get rotary embeddings
|
||||
freqs_cos, freqs_sin = get_rotary_pos_embed((tt * get_sequence_model_parallel_world_size(), th, tw), self.hidden_size, self.num_attention_heads, self.rope_dim_list, self.rope_theta)
|
||||
freqs_cos = freqs_cos.to(x.device)
|
||||
freqs_sin = freqs_sin.to(x.device)
|
||||
# Prepare modulation vectors
|
||||
vec = self.time_in(t)
|
||||
|
||||
# Add text modulation
|
||||
vec = vec + self.vector_in(text_states_2)
|
||||
|
||||
# Add guidance modulation if needed
|
||||
if self.guidance_embeds and guidance is not None:
|
||||
vec = vec + self.guidance_in(guidance)
|
||||
# Embed image and text
|
||||
img = self.img_in(img)
|
||||
txt = self.txt_in(txt, t)
|
||||
txt_seq_len = txt.shape[1]
|
||||
img_seq_len = img.shape[1]
|
||||
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
# Process through double stream blocks
|
||||
for index, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis]
|
||||
img, txt = block(*double_block_args)
|
||||
# Merge txt and img to pass through single stream blocks
|
||||
x = torch.cat((img, txt), 1)
|
||||
|
||||
|
||||
# Process through single stream blocks
|
||||
if len(self.single_blocks) > 0:
|
||||
for index, block in enumerate(self.single_blocks):
|
||||
single_block_args = [
|
||||
x,
|
||||
vec,
|
||||
txt_seq_len,
|
||||
freqs_cis,
|
||||
]
|
||||
x = block(*single_block_args)
|
||||
|
||||
# Extract image features
|
||||
img = x[:, :img_seq_len, ...]
|
||||
# Final layer processing
|
||||
img = self.final_layer(img, vec)
|
||||
# Unpatchify to get original shape
|
||||
img = unpatchify(img, tt, th, tw, self.patch_size, self.out_channels)
|
||||
|
||||
|
||||
return img
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class SingleTokenRefiner(nn.Module):
|
||||
"""
|
||||
A token refiner that processes text embeddings with attention to improve
|
||||
their representation for cross-attention with image features.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
depth=2,
|
||||
qkv_bias=True,
|
||||
dtype=None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# Input projection
|
||||
self.input_embedder = ReplicatedLinear(
|
||||
in_channels,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
# Timestep embedding
|
||||
self.t_embedder = TimestepEmbedder(
|
||||
hidden_size,
|
||||
act_layer="silu",
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
# Context embedding
|
||||
self.c_embedder = MLP(
|
||||
in_channels,
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
act_type="silu",
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
# Refiner blocks
|
||||
self.refiner_blocks = nn.ModuleList([
|
||||
IndividualTokenRefinerBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
qkv_bias=qkv_bias,
|
||||
dtype=dtype,
|
||||
) for _ in range(depth)
|
||||
])
|
||||
|
||||
def forward(self, x, t):
|
||||
# Get timestep embeddings
|
||||
timestep_aware_representations = self.t_embedder(t)
|
||||
|
||||
# Get context-aware representations
|
||||
|
||||
context_aware_representations = torch.mean(x, dim=1)
|
||||
|
||||
context_aware_representations = self.c_embedder(context_aware_representations)
|
||||
c = timestep_aware_representations + context_aware_representations
|
||||
# Project input
|
||||
x, _ = self.input_embedder(x)
|
||||
# Process through refiner blocks
|
||||
for block in self.refiner_blocks:
|
||||
x = block(x, c)
|
||||
return x
|
||||
|
||||
|
||||
class IndividualTokenRefinerBlock(nn.Module):
|
||||
"""
|
||||
A transformer block for refining individual tokens with self-attention.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
mlp_ratio=4.0,
|
||||
qkv_bias=True,
|
||||
dtype=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
|
||||
# Normalization and attention
|
||||
self.norm1 = nn.LayerNorm(hidden_size, eps=1e-6, elementwise_affine=True, dtype=dtype)
|
||||
|
||||
# self.self_attn_qkv = ReplicatedLinear(
|
||||
# hidden_size,
|
||||
# hidden_size * 3,
|
||||
# bias=qkv_bias,
|
||||
# params_dtype=dtype
|
||||
# )
|
||||
for name in ["q", "k", "v"]:
|
||||
setattr(self, f"self_attn_to_{name}", ReplicatedLinear(
|
||||
hidden_size, hidden_size, bias=True, params_dtype=dtype
|
||||
))
|
||||
|
||||
self.self_attn_proj = ReplicatedLinear(
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
bias=qkv_bias,
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
# MLP
|
||||
self.norm2 = nn.LayerNorm(hidden_size, eps=1e-6, elementwise_affine=True, dtype=dtype)
|
||||
self.mlp = MLP(
|
||||
hidden_size,
|
||||
mlp_hidden_dim,
|
||||
bias=True,
|
||||
act_type="silu",
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
# Modulation
|
||||
self.adaLN_modulation = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=2,
|
||||
act_layer="silu",
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
# Scaled dot product attention
|
||||
self.attn = LocalAttention()
|
||||
|
||||
def forward(self, x, c):
|
||||
# Get modulation parameters
|
||||
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=-1)
|
||||
# Self-attention
|
||||
norm_x = self.norm1(x)
|
||||
q_, _ = self.self_attn_to_q(norm_x)
|
||||
k_, _ = self.self_attn_to_k(norm_x)
|
||||
v_, _ = self.self_attn_to_v(norm_x)
|
||||
qkv = torch.cat((q_, k_, v_), dim=-1)
|
||||
|
||||
batch_size, seq_len = qkv.shape[0], qkv.shape[1]
|
||||
qkv = qkv.view(batch_size, seq_len, 3, self.num_attention_heads, -1)
|
||||
q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2]
|
||||
|
||||
|
||||
# Run scaled dot product attention
|
||||
attn_output = self.attn(q, k, v) # [B, L, H, D]
|
||||
attn_output = attn_output.reshape(batch_size, seq_len, -1) # [B, L, H*D]
|
||||
|
||||
# Project and apply residual connection with gating
|
||||
attn_out, _ = self.self_attn_proj(attn_output)
|
||||
x = x + attn_out * gate_msa.unsqueeze(1)
|
||||
|
||||
# MLP
|
||||
mlp_out = self.mlp(self.norm2(x))
|
||||
x = x + mlp_out * gate_mlp.unsqueeze(1)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
|
||||
|
||||
class FinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of DiT that projects features to pixel space.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, patch_size, out_channels, dtype=None):
|
||||
super().__init__()
|
||||
|
||||
# Normalization
|
||||
self.norm_final = nn.LayerNorm(hidden_size, eps=1e-6, elementwise_affine=False, dtype=dtype)
|
||||
|
||||
|
||||
output_dim = patch_size[0] * patch_size[1] * patch_size[2] * out_channels
|
||||
|
||||
self.linear = ReplicatedLinear(
|
||||
hidden_size,
|
||||
output_dim,
|
||||
bias=True,
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
|
||||
# Modulation
|
||||
self.adaLN_modulation = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=2,
|
||||
act_layer="silu",
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
def forward(self, x, c):
|
||||
# What the heck HF? Why you change the scale and shift order here???
|
||||
scale, shift = self.adaLN_modulation(c).chunk(2, dim=-1)
|
||||
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
x, _ = self.linear(x)
|
||||
return x
|
||||
@@ -0,0 +1,456 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import math
|
||||
from typing import Optional, Tuple, List, Union, Dict, Any
|
||||
from fastvideo.v1.attention.flash_attn import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
from fastvideo.v1.layers.layernorm import LayerNormScaleShift, ScaleResidual, ScaleResidualLayerNormScaleShift, RMSNorm
|
||||
from fastvideo.v1.layers.visual_embedding import PatchEmbed, TimestepEmbedder, ModulateProjection
|
||||
from fastvideo.v1.layers.rotary_embedding import _apply_rotary_emb, get_rotary_pos_embed
|
||||
from fastvideo.v1.distributed.parallel_state import get_sequence_model_parallel_world_size
|
||||
# from torch.nn import RMSNorm
|
||||
# TODO: RMSNorm ....
|
||||
from fastvideo.v1.layers.mlp import MLP
|
||||
from fastvideo.v1.models.dits.base import BaseDiT
|
||||
|
||||
class WanImageEmbedding(torch.nn.Module):
|
||||
def __init__(self, in_features: int, out_features: int):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = nn.LayerNorm(in_features)
|
||||
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
|
||||
self.norm2 = nn.LayerNorm(out_features)
|
||||
|
||||
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.norm1(encoder_hidden_states_image)
|
||||
hidden_states = self.ff(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
class WanTimeTextImageEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
time_freq_dim: int,
|
||||
text_embed_dim: int,
|
||||
image_embed_dim: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.time_embedder = TimestepEmbedder(dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
|
||||
self.time_modulation = ModulateProjection(dim, factor=6, act_layer="silu")
|
||||
self.text_embedder = MLP(text_embed_dim, dim, dim, bias=True, act_type="gelu_pytorch_tanh")
|
||||
|
||||
self.image_embedder = None
|
||||
if image_embed_dim is not None:
|
||||
self.image_embedder = WanImageEmbedding(image_embed_dim, dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_image: Optional[torch.Tensor] = None,
|
||||
):
|
||||
with torch.cuda.amp.autocast(dtype=torch.float32):
|
||||
temb = self.time_embedder(timestep.float())
|
||||
timestep_proj = self.time_modulation(temb)
|
||||
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image)
|
||||
|
||||
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
|
||||
|
||||
class WanSelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
num_heads,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
parallel_attention=False):
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.window_size = window_size
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
self.parallel_attention = parallel_attention
|
||||
|
||||
# layers
|
||||
self.to_q = ReplicatedLinear(dim, dim)
|
||||
self.to_k = ReplicatedLinear(dim, dim)
|
||||
self.to_v = ReplicatedLinear(dim, dim)
|
||||
self.to_out = ReplicatedLinear(dim, dim)
|
||||
self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
# Scaled dot product attention
|
||||
self.attn = LocalAttention(dropout_rate=0, softmax_scale=None, causal=False)
|
||||
|
||||
def forward(self, x, seq_lens, grid_sizes, freqs):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, num_heads, C / num_heads]
|
||||
seq_lens(Tensor): Shape [B]
|
||||
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
|
||||
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class WanT2VCrossAttention(WanSelfAttention):
|
||||
|
||||
def forward(self, x, context, context_lens):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
context(Tensor): Shape [B, L2, C]
|
||||
context_lens(Tensor): Shape [B]
|
||||
"""
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
|
||||
# compute attention
|
||||
x = self.attn(q, k, v)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
x, _ = self.to_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
num_heads,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6):
|
||||
super().__init__(dim, num_heads, window_size, qk_norm, eps)
|
||||
|
||||
self.add_k_proj = ReplicatedLinear(dim, dim)
|
||||
self.add_v_proj = ReplicatedLinear(dim, dim)
|
||||
self.norm_added_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_added_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def forward(self, x, context, context_lens):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
context(Tensor): Shape [B, L2, C]
|
||||
context_lens(Tensor): Shape [B]
|
||||
"""
|
||||
context_img = context[:, :257]
|
||||
context = context[:, 257:]
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
k_img = self.norm_added_k.forward_native(self.add_k_proj(context_img)[0]).view(b, -1, n, d)
|
||||
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
|
||||
img_x = self.attn(q, k_img, v_img)
|
||||
# compute attention
|
||||
x = self.attn(q, k, v)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
img_x = img_x.flatten(2)
|
||||
x = x + img_x
|
||||
x, _ = self.to_out(x)
|
||||
return x
|
||||
|
||||
class WanTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_out = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.attn1 = DistributedAttention(
|
||||
dropout_rate=0.0,
|
||||
causal=False
|
||||
)
|
||||
self.hidden_dim = dim
|
||||
self.num_attention_heads = num_heads
|
||||
dim_head = dim // num_heads
|
||||
if qk_norm == "rms_norm":
|
||||
self.norm_q = RMSNorm(dim_head, eps=eps)
|
||||
self.norm_k = RMSNorm(dim_head, eps=eps)
|
||||
elif qk_norm == "rms_norm_across_heads":
|
||||
# LTX applies qk norm across all heads
|
||||
self.norm_q = RMSNorm(dim, eps=eps)
|
||||
self.norm_k = RMSNorm(dim, eps=eps)
|
||||
else:
|
||||
print("QK Norm type not supported")
|
||||
raise Exception
|
||||
assert cross_attn_norm is True
|
||||
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(dim, norm_type="layer", eps=eps, elementwise_affine=True, dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
# I2V
|
||||
self.attn2 = WanI2VCrossAttention(dim, num_heads, qk_norm=qk_norm, eps=eps)
|
||||
else:
|
||||
# T2V
|
||||
self.attn2 = WanT2VCrossAttention(dim, num_heads, qk_norm=qk_norm, eps=eps)
|
||||
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(dim, norm_type="layer", eps=eps, elementwise_affine=False, dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
self.mlp_residual = ScaleResidual()
|
||||
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
assert temb.dtype == torch.float32
|
||||
with torch.cuda.amp.autocast(dtype=torch.float32):
|
||||
e = self.scale_shift_table + temb
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = self.norm1(hidden_states.float()).to(dtype=orig_dtype) * (1 + scale_msa) + shift_msa
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q.forward_native(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
# Apply rotary embeddings
|
||||
cos, sin = freqs_cis
|
||||
query, key = _apply_rotary_emb(query, cos, sin, is_neox_style=False), _apply_rotary_emb(key, cos, sin, is_neox_style=False)
|
||||
|
||||
attn_output, _ = self.attn1(query, key, value)
|
||||
attn_output = attn_output.flatten(2)
|
||||
attn_output, _ = self.to_out(attn_output)
|
||||
attn_output = attn_output.squeeze(1)
|
||||
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states, context=encoder_hidden_states, context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
|
||||
return hidden_states
|
||||
|
||||
class WanTransformer3DModel(BaseDiT):
|
||||
_fsdp_shard_conditions = [
|
||||
lambda n, m: "blocks" in n and str.isdigit(n.split(".")[-1]),
|
||||
]
|
||||
_param_names_mapping = {
|
||||
r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1",
|
||||
|
||||
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$": r"condition_embedder.text_embedder.fc_in.\1",
|
||||
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
r"^condition_embedder\.time_proj\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$": r"condition_embedder.image_embedder.ff.fc_in.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$": r"condition_embedder.image_embedder.ff.fc_out.\1",
|
||||
|
||||
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"blocks.\1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$": r"blocks.\1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$": r"blocks.\1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$": r"blocks.\1.to_out.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
|
||||
|
||||
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$": r"blocks.\1.attn2.to_out.\2",
|
||||
|
||||
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
|
||||
|
||||
r"blocks\.(\d+)\.norm2\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
}
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: Tuple[int] = (1, 2, 2),
|
||||
text_len = 512,
|
||||
num_attention_heads: int = 40,
|
||||
attention_head_dim: int = 128,
|
||||
in_channels: int = 16,
|
||||
out_channels: int = 16,
|
||||
text_dim: int = 4096,
|
||||
freq_dim: int = 256,
|
||||
ffn_dim: int = 13824,
|
||||
num_layers: int = 40,
|
||||
cross_attn_norm: bool = True,
|
||||
qk_norm: Optional[str] = "rms_norm_across_heads",
|
||||
eps: float = 1e-6,
|
||||
image_dim: Optional[int] = None,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
rope_max_seq_len: int = 1024,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
self.inner_dim = inner_dim
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.out_channels = out_channels or in_channels
|
||||
self.patch_size = patch_size
|
||||
self.text_len = text_len
|
||||
|
||||
# 1. Patch & position embedding
|
||||
self.patch_embedding = PatchEmbed(in_chans=in_channels, embed_dim=inner_dim, patch_size=patch_size, flatten=False)
|
||||
|
||||
# 2. Condition embeddings
|
||||
self.condition_embedder = WanTimeTextImageEmbedding(
|
||||
dim=inner_dim,
|
||||
time_freq_dim=freq_dim,
|
||||
text_embed_dim=text_dim,
|
||||
image_embed_dim=image_dim,
|
||||
)
|
||||
|
||||
# 3. Transformer blocks
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
WanTransformerBlock(
|
||||
inner_dim, ffn_dim, num_attention_heads, qk_norm, cross_attn_norm, eps, added_kv_proj_dim
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Output norm & projection
|
||||
self.norm_out = LayerNormScaleShift(inner_dim, norm_type="layer", eps=eps, elementwise_affine=False, dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size))
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
seq_len: Optional[int] = None,
|
||||
encoder_hidden_states_image: Optional[torch.Tensor] = None,
|
||||
y: Optional[torch.Tensor] = None,
|
||||
return_dict: bool = True,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
if y is not None:
|
||||
hidden_states = torch.cat([hidden_states, y], dim=1)
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
# Get rotary embeddings
|
||||
d = self.inner_dim // self.num_attention_heads
|
||||
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
|
||||
freqs_cos, freqs_sin = get_rotary_pos_embed((post_patch_num_frames * get_sequence_model_parallel_world_size(), post_patch_height, post_patch_width), self.inner_dim, self.num_attention_heads, rope_dim_list, rope_theta=10000)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
if seq_len is None:
|
||||
seq_len = hidden_states.size(1)
|
||||
hidden_states = torch.cat([hidden_states, hidden_states.new_zeros(1, seq_len - hidden_states.size(1), hidden_states.size(2))], dim=1)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image
|
||||
)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block, hidden_states, encoder_hidden_states, timestep_proj, freqs_cis
|
||||
)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, freqs_cis)
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
with torch.cuda.amp.autocast(dtype=torch.float32):
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
|
||||
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
|
||||
return output.float()
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
r"""
|
||||
Reconstruct video tensors from patch embeddings.
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
||||
grid_sizes (Tensor):
|
||||
Original spatial-temporal grid dimensions before patching,
|
||||
shape [B, 3] (3 dimensions correspond to F_patches, H_patches, W_patches)
|
||||
|
||||
Returns:
|
||||
Tensor:
|
||||
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
|
||||
c = self.out_channels
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = u.permute(6, 0, 3, 1, 4, 2, 5)
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
out = torch.cat(out, dim=0)
|
||||
return out
|
||||
@@ -0,0 +1,656 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Minimal implementation of CLIPVisionModel intended to be only used
|
||||
within a vision language model."""
|
||||
from typing import Iterable, Optional, Set, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers import CLIPVisionConfig, CLIPTextConfig
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
||||
# from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask
|
||||
|
||||
from vllm.attention.layer import MultiHeadAttention
|
||||
# from fastvideo.v1.attention.flash_attn import LocalAttention
|
||||
from fastvideo.v1.distributed import divide, get_tensor_model_parallel_world_size
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.layers.linear import (ColumnParallelLinear,
|
||||
QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
# TODO: support quantization
|
||||
# from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
|
||||
from vllm.model_executor.models.interfaces import SupportsQuant
|
||||
|
||||
from .vision import VisionEncoderInfo, resolve_visual_encoder_outputs
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
class QuantizationConfig:
|
||||
pass
|
||||
|
||||
class CLIPEncoderInfo(VisionEncoderInfo[CLIPVisionConfig]):
|
||||
|
||||
def get_num_image_tokens(
|
||||
self,
|
||||
*,
|
||||
image_width: int,
|
||||
image_height: int,
|
||||
) -> int:
|
||||
return self.get_patch_grid_length()**2 + 1
|
||||
|
||||
def get_max_image_tokens(self) -> int:
|
||||
return self.get_patch_grid_length()**2 + 1
|
||||
|
||||
def get_image_size(self) -> int:
|
||||
return self.vision_config.image_size
|
||||
|
||||
def get_patch_size(self) -> int:
|
||||
return self.vision_config.patch_size
|
||||
|
||||
def get_patch_grid_length(self) -> int:
|
||||
image_size, patch_size = self.get_image_size(), self.get_patch_size()
|
||||
assert image_size % patch_size == 0
|
||||
return image_size // patch_size
|
||||
|
||||
|
||||
# Adapted from https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py#L164 # noqa
|
||||
class CLIPVisionEmbeddings(nn.Module):
|
||||
|
||||
def __init__(self, config: CLIPVisionConfig):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.embed_dim = config.hidden_size
|
||||
self.image_size = config.image_size
|
||||
self.patch_size = config.patch_size
|
||||
assert self.image_size % self.patch_size == 0
|
||||
|
||||
self.class_embedding = nn.Parameter(torch.randn(self.embed_dim))
|
||||
|
||||
self.patch_embedding = nn.Conv2d(
|
||||
in_channels=config.num_channels,
|
||||
out_channels=self.embed_dim,
|
||||
kernel_size=self.patch_size,
|
||||
stride=self.patch_size,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.num_patches = (self.image_size // self.patch_size)**2
|
||||
self.num_positions = self.num_patches + 1
|
||||
self.position_embedding = nn.Embedding(self.num_positions,
|
||||
self.embed_dim)
|
||||
self.register_buffer("position_ids",
|
||||
torch.arange(self.num_positions).expand((1, -1)),
|
||||
persistent=False)
|
||||
|
||||
def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:
|
||||
batch_size = pixel_values.shape[0]
|
||||
target_dtype = self.patch_embedding.weight.dtype
|
||||
patch_embeds = self.patch_embedding(pixel_values.to(
|
||||
dtype=target_dtype)) # shape = [*, width, grid, grid]
|
||||
patch_embeds = patch_embeds.flatten(2).transpose(1, 2)
|
||||
|
||||
class_embeds = self.class_embedding.expand(batch_size, 1, -1)
|
||||
embeddings = torch.cat([class_embeds, patch_embeds], dim=1)
|
||||
embeddings = embeddings + self.position_embedding(self.position_ids)
|
||||
|
||||
return embeddings
|
||||
|
||||
|
||||
class CLIPTextEmbeddings(nn.Module):
|
||||
|
||||
def __init__(self, config: CLIPTextConfig):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
embed_dim = config.hidden_size
|
||||
|
||||
self.token_embedding = nn.Embedding(config.vocab_size, embed_dim)
|
||||
self.position_embedding = nn.Embedding(config.max_position_embeddings, embed_dim)
|
||||
|
||||
# position_ids (1, len position emb) is contiguous in memory and exported when serialized
|
||||
self.register_buffer(
|
||||
"position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)), persistent=False
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.LongTensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.Tensor:
|
||||
seq_length = input_ids.shape[-1] if input_ids is not None else inputs_embeds.shape[-2]
|
||||
max_position_embedding = self.position_embedding.weight.shape[0]
|
||||
|
||||
if seq_length > max_position_embedding:
|
||||
raise ValueError(
|
||||
f"Sequence length must be less than max_position_embeddings (got `sequence length`: "
|
||||
f"{seq_length} and max_position_embeddings: {max_position_embedding}"
|
||||
)
|
||||
|
||||
if position_ids is None:
|
||||
position_ids = self.position_ids[:, :seq_length]
|
||||
|
||||
if inputs_embeds is None:
|
||||
inputs_embeds = self.token_embedding(input_ids)
|
||||
|
||||
position_embeddings = self.position_embedding(position_ids)
|
||||
embeddings = inputs_embeds + position_embeddings
|
||||
|
||||
return embeddings
|
||||
|
||||
|
||||
class CLIPAttention(nn.Module):
|
||||
"""Multi-headed attention from 'Attention Is All You Need' paper"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.embed_dim = config.hidden_size
|
||||
self.num_heads = config.num_attention_heads
|
||||
self.head_dim = self.embed_dim // self.num_heads
|
||||
if self.head_dim * self.num_heads != self.embed_dim:
|
||||
raise ValueError(
|
||||
"embed_dim must be divisible by num_heads "
|
||||
f"(got `embed_dim`: {self.embed_dim} and `num_heads`:"
|
||||
f" {self.num_heads}).")
|
||||
self.scale = self.head_dim**-0.5
|
||||
self.dropout = config.attention_dropout
|
||||
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
hidden_size=self.embed_dim,
|
||||
head_size=self.head_dim,
|
||||
total_num_heads=self.num_heads,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.qkv_proj",
|
||||
)
|
||||
|
||||
self.out_proj = RowParallelLinear(
|
||||
input_size=self.embed_dim,
|
||||
output_size=self.embed_dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.out_proj",
|
||||
)
|
||||
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
self.num_heads_per_partition = divide(self.num_heads, self.tp_size)
|
||||
|
||||
self.attn = MultiHeadAttention(self.num_heads_per_partition,
|
||||
self.head_dim, self.scale)
|
||||
|
||||
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
|
||||
return tensor.view(bsz, seq_len, self.num_heads,
|
||||
self.head_dim).transpose(1, 2).contiguous()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
):
|
||||
"""Input shape: Batch x Time x Channel"""
|
||||
|
||||
qkv_states, _ = self.qkv_proj(hidden_states)
|
||||
query_states, key_states, value_states = qkv_states.chunk(3, dim=-1)
|
||||
# use flash_attn_func
|
||||
from flash_attn import flash_attn_func
|
||||
query_states = query_states.reshape(query_states.shape[0], query_states.shape[1], self.num_heads_per_partition, self.head_dim)
|
||||
key_states = key_states.reshape(key_states.shape[0], key_states.shape[1], self.num_heads_per_partition, self.head_dim)
|
||||
value_states = value_states.reshape(value_states.shape[0], value_states.shape[1], self.num_heads_per_partition, self.head_dim)
|
||||
attn_output = flash_attn_func(query_states, key_states, value_states,softmax_scale=self.scale, causal=True)
|
||||
attn_output = attn_output.reshape(attn_output.shape[0], attn_output.shape[1], self.num_heads_per_partition * self.head_dim)
|
||||
attn_output, _ = self.out_proj(attn_output)
|
||||
|
||||
return attn_output, None
|
||||
|
||||
|
||||
class CLIPMLP(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.activation_fn = get_act_fn(config.hidden_act)
|
||||
self.fc1 = ColumnParallelLinear(config.hidden_size,
|
||||
config.intermediate_size,
|
||||
bias=True,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.fc1")
|
||||
self.fc2 = RowParallelLinear(config.intermediate_size,
|
||||
config.hidden_size,
|
||||
bias=True,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.fc2")
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states, _ = self.fc1(hidden_states)
|
||||
hidden_states = self.activation_fn(hidden_states)
|
||||
hidden_states, _ = self.fc2(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class CLIPEncoderLayer(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPTextConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.self_attn = CLIPAttention(
|
||||
config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.self_attn",
|
||||
)
|
||||
self.layer_norm1 = nn.LayerNorm(config.hidden_size,
|
||||
eps=config.layer_norm_eps)
|
||||
self.mlp = CLIPMLP(config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.mlp")
|
||||
self.layer_norm2 = nn.LayerNorm(config.hidden_size,
|
||||
eps=config.layer_norm_eps)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.layer_norm1(hidden_states)
|
||||
hidden_states, _ = self.self_attn(hidden_states=hidden_states)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
residual = hidden_states
|
||||
hidden_states = self.layer_norm2(hidden_states)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class CLIPEncoder(nn.Module):
|
||||
"""
|
||||
Transformer encoder consisting of `config.num_hidden_layers` self
|
||||
attention layers. Each layer is a [`CLIPEncoderLayer`].
|
||||
|
||||
Args:
|
||||
config: CLIPConfig
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.config = config
|
||||
|
||||
if num_hidden_layers_override is None:
|
||||
num_hidden_layers = config.num_hidden_layers
|
||||
else:
|
||||
num_hidden_layers = num_hidden_layers_override
|
||||
self.layers = nn.ModuleList([
|
||||
CLIPEncoderLayer(config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.layers.{layer_idx}")
|
||||
for layer_idx in range(num_hidden_layers)
|
||||
])
|
||||
|
||||
def forward(
|
||||
self, inputs_embeds: torch.Tensor, return_all_hidden_states: bool
|
||||
) -> Union[torch.Tensor, list[torch.Tensor]]:
|
||||
hidden_states_pool = [inputs_embeds]
|
||||
hidden_states = inputs_embeds
|
||||
|
||||
for encoder_layer in self.layers:
|
||||
hidden_states = encoder_layer(hidden_states)
|
||||
if return_all_hidden_states:
|
||||
hidden_states_pool.append(hidden_states)
|
||||
# If we have multiple feature sample layers, we return all hidden
|
||||
# states in order and grab the ones we need by index.
|
||||
if return_all_hidden_states:
|
||||
return hidden_states_pool
|
||||
return [hidden_states]
|
||||
|
||||
|
||||
|
||||
class CLIPTextTransformer(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: CLIPTextConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
*,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
embed_dim = config.hidden_size
|
||||
|
||||
self.embeddings = CLIPTextEmbeddings(config)
|
||||
|
||||
self.encoder = CLIPEncoder(config,
|
||||
quant_config=quant_config,
|
||||
num_hidden_layers_override=num_hidden_layers_override,
|
||||
prefix=prefix)
|
||||
|
||||
self.final_layer_norm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
|
||||
|
||||
# For `pooled_output` computation
|
||||
self.eos_token_id = config.eos_token_id
|
||||
|
||||
# For attention mask, it differs between `flash_attention_2` and other attention implementations
|
||||
self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2"
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = None,
|
||||
) -> Union[Tuple, BaseModelOutputWithPooling]:
|
||||
r"""
|
||||
Returns:
|
||||
|
||||
"""
|
||||
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
|
||||
output_hidden_states = (
|
||||
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
||||
)
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
if input_ids is None:
|
||||
raise ValueError("You have to specify input_ids")
|
||||
|
||||
input_shape = input_ids.size()
|
||||
input_ids = input_ids.view(-1, input_shape[-1])
|
||||
|
||||
hidden_states = self.embeddings(input_ids=input_ids, position_ids=position_ids)
|
||||
|
||||
# CLIP's text model uses causal mask, prepare it here.
|
||||
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
|
||||
# causal_attention_mask = _create_4d_causal_attention_mask(
|
||||
# input_shape, hidden_states.dtype, device=hidden_states.device
|
||||
# )
|
||||
|
||||
# # expand attention_mask
|
||||
# if attention_mask is not None and not self._use_flash_attention_2:
|
||||
# raise NotImplementedError("attention_mask is not supported for CLIPTextTransformer")
|
||||
# # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
|
||||
# attention_mask = _prepare_4d_attention_mask(attention_mask, hidden_states.dtype)
|
||||
|
||||
encoder_outputs = self.encoder(
|
||||
inputs_embeds=hidden_states,
|
||||
# attention_mask=attention_mask,
|
||||
# causal_attention_mask=causal_attention_mask,
|
||||
# output_attentions=output_attentions,
|
||||
return_all_hidden_states=output_hidden_states,
|
||||
# return_dict=return_dict,
|
||||
)
|
||||
|
||||
last_hidden_state = encoder_outputs[-1]
|
||||
last_hidden_state = self.final_layer_norm(last_hidden_state)
|
||||
|
||||
if self.eos_token_id == 2:
|
||||
# The `eos_token_id` was incorrect before PR #24773: Let's keep what have been done here.
|
||||
# A CLIP model with such `eos_token_id` in the config can't work correctly with extra new tokens added
|
||||
# ------------------------------------------------------------
|
||||
# text_embeds.shape = [batch_size, sequence_length, transformer.width]
|
||||
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||
# casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14
|
||||
pooled_output = last_hidden_state[
|
||||
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
|
||||
input_ids.to(dtype=torch.int, device=last_hidden_state.device).argmax(dim=-1),
|
||||
]
|
||||
else:
|
||||
# The config gets updated `eos_token_id` from PR #24773 (so the use of exta new tokens is possible)
|
||||
pooled_output = last_hidden_state[
|
||||
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
|
||||
# We need to get the first position of `eos_token_id` value (`pad_token_ids` might equal to `eos_token_id`)
|
||||
# Note: we assume each sequence (along batch dim.) contains an `eos_token_id` (e.g. prepared by the tokenizer)
|
||||
(input_ids.to(dtype=torch.int, device=last_hidden_state.device) == self.eos_token_id)
|
||||
.int()
|
||||
.argmax(dim=-1),
|
||||
]
|
||||
|
||||
if not return_dict:
|
||||
return (last_hidden_state, pooled_output) + encoder_outputs[1:]
|
||||
|
||||
# return last_hidden_state
|
||||
return BaseModelOutputWithPooling(
|
||||
last_hidden_state=last_hidden_state,
|
||||
pooler_output=pooled_output,
|
||||
hidden_states=encoder_outputs,
|
||||
# attentions=encoder_outputs.attentions,
|
||||
)
|
||||
|
||||
|
||||
class CLIPTextModel(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPTextConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.config = config
|
||||
self.text_model = CLIPTextTransformer(
|
||||
config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = None,
|
||||
) -> Union[Tuple, BaseModelOutputWithPooling]:
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
return self.text_model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
position_ids=position_ids,
|
||||
output_attentions=output_attentions,
|
||||
output_hidden_states=output_hidden_states,
|
||||
return_dict=None,
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
|
||||
# Define mapping for stacked parameters
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
# Handle q_proj, k_proj, v_proj -> qkv_proj mapping
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
if weight_name in name:
|
||||
# Replace the weight name with the parameter name
|
||||
model_param_name = name.replace(weight_name, param_name)
|
||||
|
||||
if model_param_name in params_dict:
|
||||
param = params_dict[model_param_name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
loaded_params.add(model_param_name)
|
||||
break
|
||||
else:
|
||||
# Use default weight loader for all other parameters
|
||||
if name in params_dict:
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
|
||||
return loaded_params
|
||||
|
||||
|
||||
class CLIPVisionTransformer(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
*,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
require_post_norm: Optional[bool] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.config = config
|
||||
embed_dim = config.hidden_size
|
||||
|
||||
self.embeddings = CLIPVisionEmbeddings(config)
|
||||
|
||||
# NOTE: This typo of "layrnorm" is not fixed on purpose to match
|
||||
# the original transformers code and name of the model weights.
|
||||
self.pre_layrnorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
|
||||
|
||||
self.encoder = CLIPEncoder(
|
||||
config=config,
|
||||
quant_config=quant_config,
|
||||
num_hidden_layers_override=num_hidden_layers_override,
|
||||
prefix=f"{prefix}.encoder",
|
||||
)
|
||||
|
||||
num_hidden_layers = config.num_hidden_layers
|
||||
if len(self.encoder.layers) > config.num_hidden_layers:
|
||||
raise ValueError(
|
||||
f"The original encoder only has {num_hidden_layers} "
|
||||
f"layers, but you requested {len(self.encoder.layers)} layers."
|
||||
)
|
||||
|
||||
# If possible, skip post_layernorm to conserve memory
|
||||
if require_post_norm is None:
|
||||
require_post_norm = len(self.encoder.layers) == num_hidden_layers
|
||||
|
||||
if require_post_norm:
|
||||
self.post_layernorm = nn.LayerNorm(embed_dim,
|
||||
eps=config.layer_norm_eps)
|
||||
else:
|
||||
self.post_layernorm = None
|
||||
|
||||
def forward(
|
||||
self,
|
||||
pixel_values: torch.Tensor,
|
||||
feature_sample_layers: Optional[list[int]] = None,
|
||||
) -> torch.Tensor:
|
||||
|
||||
hidden_states = self.embeddings(pixel_values)
|
||||
hidden_states = self.pre_layrnorm(hidden_states)
|
||||
|
||||
return_all_hidden_states = feature_sample_layers is not None
|
||||
|
||||
# Produces either the last layer output or all of the hidden states,
|
||||
# depending on if we have feature_sample_layers or not
|
||||
encoder_outputs = self.encoder(
|
||||
inputs_embeds=hidden_states,
|
||||
return_all_hidden_states=return_all_hidden_states)
|
||||
|
||||
# Handle post-norm (if applicable) and stacks feature layers if needed
|
||||
encoder_outputs = resolve_visual_encoder_outputs(
|
||||
encoder_outputs, feature_sample_layers, self.post_layernorm,
|
||||
self.config.num_hidden_layers)
|
||||
|
||||
return encoder_outputs
|
||||
|
||||
|
||||
class CLIPVisionModel(nn.Module, SupportsQuant):
|
||||
config_class = CLIPVisionConfig
|
||||
main_input_name = "pixel_values"
|
||||
packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
*,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
require_post_norm: Optional[bool] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.vision_model = CLIPVisionTransformer(
|
||||
config=config,
|
||||
quant_config=quant_config,
|
||||
num_hidden_layers_override=num_hidden_layers_override,
|
||||
require_post_norm=require_post_norm,
|
||||
prefix=f"{prefix}.vision_model")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
pixel_values: torch.Tensor,
|
||||
feature_sample_layers: Optional[list[int]] = None,
|
||||
) -> torch.Tensor:
|
||||
return self.vision_model(pixel_values, feature_sample_layers)
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
# (TODO) Add prefix argument for filtering out weights to be loaded
|
||||
# ref: https://github.com/vllm-project/vllm/pull/7186#discussion_r1734163986
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
layer_count = len(self.vision_model.encoder.layers)
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
# post_layernorm is not needed in CLIPVisionModel
|
||||
if (name.startswith("vision_model.post_layernorm")
|
||||
and self.vision_model.post_layernorm is None):
|
||||
continue
|
||||
|
||||
# omit layers when num_hidden_layers_override is set
|
||||
if name.startswith("vision_model.encoder.layers"):
|
||||
layer_idx = int(name.split(".")[3])
|
||||
if layer_idx >= layer_count:
|
||||
continue
|
||||
|
||||
for (param_name, weight_name, shard_id) in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
break
|
||||
else:
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader",
|
||||
default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
return loaded_params
|
||||
@@ -0,0 +1,424 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Adapted from
|
||||
# https://github.com/huggingface/transformers/blob/v4.28.0/src/transformers/models/llama/modeling_llama.py
|
||||
# Copyright 2023 The vLLM team.
|
||||
# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
|
||||
# and OPT implementations in this library. It has been modified from its
|
||||
# original forms to accommodate minor architectural differences compared
|
||||
# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Inference-only LLaMA model compatible with HuggingFace weights."""
|
||||
from typing import Any, Dict, Iterable, Optional, Set, Tuple, Type, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from transformers import LlamaConfig
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPast
|
||||
|
||||
from vllm.attention.layer import MultiHeadAttention
|
||||
|
||||
from fastvideo.v1.distributed import get_tensor_model_parallel_world_size
|
||||
from fastvideo.v1.layers.activation import SiluAndMul
|
||||
from fastvideo.v1.layers.layernorm import RMSNorm
|
||||
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
|
||||
QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
# from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
|
||||
from fastvideo.v1.layers.rotary_embedding import get_rope
|
||||
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.v1.models.loader.weight_utils import (
|
||||
default_weight_loader, maybe_remap_kv_scale_name)
|
||||
|
||||
from .utils import (extract_layer_index)
|
||||
|
||||
class QuantizationConfig:
|
||||
pass
|
||||
|
||||
|
||||
class LlamaMLP(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
intermediate_size: int,
|
||||
hidden_act: str,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
bias: bool = False,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.gate_up_proj = MergedColumnParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_sizes=[intermediate_size] * 2,
|
||||
# output_size=intermediate_size,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.gate_up_proj",
|
||||
)
|
||||
self.down_proj = RowParallelLinear(
|
||||
input_size=intermediate_size,
|
||||
output_size=hidden_size,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.down_proj",
|
||||
)
|
||||
if hidden_act != "silu":
|
||||
raise ValueError(f"Unsupported activation: {hidden_act}. "
|
||||
"Only silu is supported for now.")
|
||||
self.act_fn = SiluAndMul()
|
||||
|
||||
def forward(self, x):
|
||||
x, _ = self.gate_up_proj(x)
|
||||
x = self.act_fn(x)
|
||||
x, _ = self.down_proj(x)
|
||||
return x
|
||||
|
||||
|
||||
class LlamaAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: LlamaConfig,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
num_kv_heads: int,
|
||||
rope_theta: float = 10000,
|
||||
rope_scaling: Optional[Dict[str, Any]] = None,
|
||||
max_position_embeddings: int = 8192,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
bias: bool = False,
|
||||
bias_o_proj: bool = False,
|
||||
prefix: str = "") -> None:
|
||||
super().__init__()
|
||||
layer_idx = extract_layer_index(prefix)
|
||||
self.hidden_size = hidden_size
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
self.total_num_heads = num_heads
|
||||
assert self.total_num_heads % tp_size == 0
|
||||
self.num_heads = self.total_num_heads // tp_size
|
||||
self.total_num_kv_heads = num_kv_heads
|
||||
if self.total_num_kv_heads >= tp_size:
|
||||
# Number of KV heads is greater than TP size, so we partition
|
||||
# the KV heads across multiple tensor parallel GPUs.
|
||||
assert self.total_num_kv_heads % tp_size == 0
|
||||
else:
|
||||
# Number of KV heads is less than TP size, so we replicate
|
||||
# the KV heads across multiple tensor parallel GPUs.
|
||||
assert tp_size % self.total_num_kv_heads == 0
|
||||
self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)
|
||||
# MistralConfig has an optional head_dim introduced by Mistral-Nemo
|
||||
self.head_dim = getattr(config, "head_dim",
|
||||
self.hidden_size // self.total_num_heads)
|
||||
# Phi models introduced a partial_rotary_factor parameter in the config
|
||||
partial_rotary_factor = getattr(config, "partial_rotary_factor", 1)
|
||||
self.rotary_dim = int(partial_rotary_factor * self.head_dim)
|
||||
self.q_size = self.num_heads * self.head_dim
|
||||
self.kv_size = self.num_kv_heads * self.head_dim
|
||||
self.scaling = self.head_dim**-0.5
|
||||
self.rope_theta = rope_theta
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
hidden_size=hidden_size,
|
||||
head_size=self.head_dim,
|
||||
total_num_heads=self.total_num_heads,
|
||||
total_num_kv_heads=self.total_num_kv_heads,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.qkv_proj",
|
||||
)
|
||||
|
||||
self.o_proj = RowParallelLinear(
|
||||
input_size=self.total_num_heads * self.head_dim,
|
||||
output_size=hidden_size,
|
||||
bias=bias_o_proj,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.o_proj",
|
||||
)
|
||||
|
||||
is_neox_style = True
|
||||
is_gguf = quant_config and quant_config.get_name() == "gguf"
|
||||
if is_gguf and config.model_type == "llama":
|
||||
is_neox_style = False
|
||||
|
||||
self.rotary_emb = get_rope(
|
||||
self.head_dim,
|
||||
rotary_dim=self.rotary_dim,
|
||||
max_position=max_position_embeddings,
|
||||
base=rope_theta,
|
||||
rope_scaling=rope_scaling,
|
||||
is_neox_style=is_neox_style,
|
||||
)
|
||||
|
||||
self.attn = MultiHeadAttention(self.num_heads,
|
||||
self.head_dim,
|
||||
self.scaling,
|
||||
self.num_kv_heads)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||
q, k = self.rotary_emb(positions, q, k)
|
||||
# attn_output = self.attn(q, k, v)
|
||||
# use flash_attn_func
|
||||
# TODO (Attn abstraction and backend)
|
||||
from flash_attn import flash_attn_func
|
||||
# reshape q, k, v to (batch_size, seq_len, num_heads, head_dim)
|
||||
batch_size = q.shape[0]
|
||||
seq_len = q.shape[1]
|
||||
q = q.reshape(batch_size, seq_len, self.num_heads, self.head_dim)
|
||||
k = k.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim)
|
||||
v = v.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim)
|
||||
# import pdb; pdb.set_trace()
|
||||
attn_output = flash_attn_func(q, k, v, softmax_scale=self.scaling, causal=True)
|
||||
attn_output = attn_output.reshape(batch_size, seq_len, self.num_heads * self.head_dim)
|
||||
|
||||
output, _ = self.o_proj(attn_output)
|
||||
return output
|
||||
|
||||
|
||||
class LlamaDecoderLayer(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: LlamaConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = config.hidden_size
|
||||
rope_theta = getattr(config, "rope_theta", 10000)
|
||||
rope_scaling = getattr(config, "rope_scaling", None)
|
||||
if rope_scaling is not None and getattr(
|
||||
config, "original_max_position_embeddings", None):
|
||||
rope_scaling["original_max_position_embeddings"] = (
|
||||
config.original_max_position_embeddings)
|
||||
max_position_embeddings = getattr(config, "max_position_embeddings",
|
||||
8192)
|
||||
# Support abacusai/Smaug-72B-v0.1 with attention_bias
|
||||
# Support internlm/internlm-7b with bias
|
||||
attention_bias = getattr(config, "attention_bias", False) or getattr(
|
||||
config, "bias", False)
|
||||
bias_o_proj = attention_bias
|
||||
# support internlm/internlm3-8b with qkv_bias
|
||||
if hasattr(config, 'qkv_bias'):
|
||||
attention_bias = config.qkv_bias
|
||||
|
||||
self.self_attn = LlamaAttention(
|
||||
config=config,
|
||||
hidden_size=self.hidden_size,
|
||||
num_heads=config.num_attention_heads,
|
||||
num_kv_heads=getattr(config, "num_key_value_heads",
|
||||
config.num_attention_heads),
|
||||
rope_theta=rope_theta,
|
||||
rope_scaling=rope_scaling,
|
||||
max_position_embeddings=max_position_embeddings,
|
||||
quant_config=quant_config,
|
||||
bias=attention_bias,
|
||||
bias_o_proj=bias_o_proj,
|
||||
prefix=f"{prefix}.self_attn",
|
||||
)
|
||||
self.mlp = LlamaMLP(
|
||||
hidden_size=self.hidden_size,
|
||||
intermediate_size=config.intermediate_size,
|
||||
hidden_act=config.hidden_act,
|
||||
quant_config=quant_config,
|
||||
bias=getattr(config, "mlp_bias", False),
|
||||
prefix=f"{prefix}.mlp",
|
||||
)
|
||||
self.input_layernorm = RMSNorm(config.hidden_size,
|
||||
eps=config.rms_norm_eps)
|
||||
self.post_attention_layernorm = RMSNorm(config.hidden_size,
|
||||
eps=config.rms_norm_eps)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
residual: Optional[torch.Tensor],
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# Self Attention
|
||||
if residual is None:
|
||||
residual = hidden_states
|
||||
hidden_states = self.input_layernorm(hidden_states)
|
||||
else:
|
||||
hidden_states, residual = self.input_layernorm(
|
||||
hidden_states, residual)
|
||||
|
||||
hidden_states = self.self_attn(positions=positions,
|
||||
hidden_states=hidden_states)
|
||||
|
||||
|
||||
# Fully Connected
|
||||
hidden_states, residual = self.post_attention_layernorm(
|
||||
hidden_states, residual)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
return hidden_states, residual
|
||||
|
||||
|
||||
class LlamaModel(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: LlamaConfig,
|
||||
prefix: str = "",
|
||||
layer_type: Type[LlamaDecoderLayer] = LlamaDecoderLayer):
|
||||
super().__init__()
|
||||
|
||||
quant_config = None
|
||||
lora_config = None
|
||||
|
||||
self.config = config
|
||||
self.quant_config = quant_config
|
||||
lora_vocab = (lora_config.lora_extra_vocab_size *
|
||||
(lora_config.max_loras or 1)) if lora_config else 0
|
||||
self.vocab_size = config.vocab_size + lora_vocab
|
||||
self.org_vocab_size = config.vocab_size
|
||||
|
||||
self.embed_tokens = VocabParallelEmbedding(
|
||||
self.vocab_size,
|
||||
config.hidden_size,
|
||||
org_num_embeddings=config.vocab_size,
|
||||
quant_config=quant_config,
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList([
|
||||
layer_type(
|
||||
config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.layers.{i}"
|
||||
) for i in range(config.num_hidden_layers)
|
||||
])
|
||||
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
|
||||
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
return self.embed_tokens(input_ids)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
positions: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
) -> torch.Tensor:
|
||||
output_hidden_states = (
|
||||
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
||||
)
|
||||
if inputs_embeds is not None:
|
||||
hidden_states = inputs_embeds
|
||||
else:
|
||||
hidden_states = self.get_input_embeddings(input_ids)
|
||||
residual = None
|
||||
|
||||
if positions is None:
|
||||
positions = torch.arange(
|
||||
0, hidden_states.shape[1], device=hidden_states.device
|
||||
).unsqueeze(0)
|
||||
|
||||
all_hidden_states = () if output_hidden_states else None
|
||||
for layer in self.layers:
|
||||
if output_hidden_states:
|
||||
all_hidden_states += (hidden_states,)
|
||||
hidden_states, residual = layer(positions, hidden_states, residual)
|
||||
|
||||
hidden_states, _ = self.norm(hidden_states, residual)
|
||||
|
||||
# add hidden states from the last decoder layer
|
||||
if output_hidden_states:
|
||||
all_hidden_states += (hidden_states,)
|
||||
|
||||
# TODO(will): maybe unify the output format with other models and use
|
||||
# our own class
|
||||
output = BaseModelOutputWithPast(
|
||||
last_hidden_state=hidden_states,
|
||||
# past_key_values=past_key_values if use_cache else None,
|
||||
hidden_states=all_hidden_states,
|
||||
# attentions=all_self_attns,
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q_proj", "q"),
|
||||
(".qkv_proj", ".k_proj", "k"),
|
||||
(".qkv_proj", ".v_proj", "v"),
|
||||
(".gate_up_proj", ".gate_proj", 0),
|
||||
(".gate_up_proj", ".up_proj", 1),
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
if "rotary_emb.inv_freq" in name:
|
||||
continue
|
||||
if ("rotary_emb.cos_cached" in name
|
||||
or "rotary_emb.sin_cached" in name):
|
||||
# Models trained using ColossalAI may include these tensors in
|
||||
# the checkpoint. Skip them.
|
||||
continue
|
||||
if (self.quant_config is not None and
|
||||
(scale_name := self.quant_config.get_cache_scale(name))):
|
||||
# Loading kv cache quantization scales
|
||||
param = params_dict[scale_name]
|
||||
weight_loader = getattr(param, "weight_loader",
|
||||
default_weight_loader)
|
||||
loaded_weight = (loaded_weight if loaded_weight.dim() == 0 else
|
||||
loaded_weight[0])
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(scale_name)
|
||||
continue
|
||||
if "scale" in name:
|
||||
# Remapping the name of FP8 kv-scale.
|
||||
name = maybe_remap_kv_scale_name(name, params_dict)
|
||||
if name is None:
|
||||
continue
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
break
|
||||
else:
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader",
|
||||
default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
return loaded_params
|
||||
@@ -0,0 +1,787 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Derived from T5 implementation posted on HuggingFace; license below:
|
||||
#
|
||||
# coding=utf-8
|
||||
# Copyright 2018 Mesh TensorFlow authors, T5 Authors and HuggingFace Inc. team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""PyTorch T5 model."""
|
||||
|
||||
import math
|
||||
import re
|
||||
from typing import Iterable, List, Optional, Set, Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from transformers import T5Config
|
||||
# TODO best way to handle xformers imports?
|
||||
from xformers.ops.fmha.attn_bias import LowerTriangularMaskWithTensorBias
|
||||
|
||||
# TODO func should be in backend interface
|
||||
from vllm.attention.backends.xformers import (XFormersMetadata, _get_attn_bias,
|
||||
_set_attn_bias)
|
||||
from vllm.attention.layer import Attention, AttentionMetadata, AttentionType
|
||||
from vllm.config import CacheConfig, VllmConfig
|
||||
from vllm.distributed import get_tensor_model_parallel_world_size
|
||||
from vllm.model_executor.layers.activation import get_act_fn
|
||||
from vllm.model_executor.layers.linear import (ColumnParallelLinear,
|
||||
QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
from vllm.model_executor.layers.logits_processor import LogitsProcessor
|
||||
from vllm.model_executor.layers.quantization.base_config import (
|
||||
QuantizationConfig)
|
||||
from vllm.model_executor.layers.sampler import SamplerOutput, get_sampler
|
||||
from vllm.model_executor.layers.vocab_parallel_embedding import (
|
||||
ParallelLMHead, VocabParallelEmbedding)
|
||||
from vllm.model_executor.model_loader.weight_utils import default_weight_loader
|
||||
from vllm.model_executor.sampling_metadata import SamplingMetadata
|
||||
from vllm.sequence import IntermediateTensors
|
||||
|
||||
from .utils import maybe_prefix
|
||||
|
||||
|
||||
class T5LayerNorm(nn.Module):
|
||||
|
||||
def __init__(self, hidden_size, eps=1e-6):
|
||||
"""
|
||||
Construct a layernorm module in the T5 style.
|
||||
No bias and no subtraction of mean.
|
||||
"""
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(hidden_size))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
def forward(self, hidden_states) -> torch.Tensor:
|
||||
# T5 uses a layer_norm which only scales and doesn't shift, which is
|
||||
# also known as Root Mean Square Layer Normalization
|
||||
# https://arxiv.org/abs/1910.07467 thus variance is calculated w/o mean
|
||||
# and there is no bias. Additionally we want to make sure that the
|
||||
# accumulation for half-precision inputs is done in fp32.
|
||||
# TODO (rmns norm ops)
|
||||
variance = hidden_states.to(torch.float32).pow(2).mean(-1,
|
||||
keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance +
|
||||
self.variance_epsilon)
|
||||
|
||||
# convert into half-precision if necessary
|
||||
if self.weight.dtype in [torch.float16, torch.bfloat16]:
|
||||
hidden_states = hidden_states.to(self.weight.dtype)
|
||||
|
||||
return self.weight * hidden_states
|
||||
|
||||
|
||||
class T5DenseActDense(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: T5Config,
|
||||
quant_config: Optional[QuantizationConfig] = None):
|
||||
super().__init__()
|
||||
self.wi = ColumnParallelLinear(config.d_model, config.d_ff, bias=False)
|
||||
self.wo = RowParallelLinear(config.d_ff,
|
||||
config.d_model,
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
self.act = get_act_fn(config.dense_act_fn)
|
||||
|
||||
def forward(self, hidden_states) -> torch.Tensor:
|
||||
hidden_states, _ = self.wi(hidden_states)
|
||||
hidden_states = self.act(hidden_states)
|
||||
# if (
|
||||
# isinstance(self.wo.weight, torch.Tensor)
|
||||
# and hidden_states.dtype != self.wo.weight.dtype
|
||||
# and self.wo.weight.dtype != torch.int8
|
||||
# ):
|
||||
# hidden_states = hidden_states.to(self.wo.weight.dtype)
|
||||
hidden_states, _ = self.wo(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5DenseGatedActDense(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: T5Config,
|
||||
quant_config: Optional[QuantizationConfig] = None):
|
||||
super().__init__()
|
||||
self.wi_0 = ColumnParallelLinear(config.d_model,
|
||||
config.d_ff,
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
self.wi_1 = ColumnParallelLinear(config.d_model,
|
||||
config.d_ff,
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
# Should not run in fp16 unless mixed-precision is used,
|
||||
# see https://github.com/huggingface/transformers/issues/20287.
|
||||
self.wo = RowParallelLinear(config.d_ff,
|
||||
config.d_model,
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
self.act = get_act_fn(config.dense_act_fn)
|
||||
|
||||
def forward(self, hidden_states) -> torch.Tensor:
|
||||
hidden_gelu = self.act(self.wi_0(hidden_states)[0])
|
||||
hidden_linear, _ = self.wi_1(hidden_states)
|
||||
hidden_states = hidden_gelu * hidden_linear
|
||||
hidden_states, _ = self.wo(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5LayerFF(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: T5Config,
|
||||
quant_config: Optional[QuantizationConfig] = None):
|
||||
super().__init__()
|
||||
if config.is_gated_act:
|
||||
self.DenseReluDense = T5DenseGatedActDense(
|
||||
config, quant_config=quant_config)
|
||||
else:
|
||||
self.DenseReluDense = T5DenseActDense(config,
|
||||
quant_config=quant_config)
|
||||
|
||||
self.layer_norm = T5LayerNorm(config.d_model,
|
||||
eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(self, hidden_states) -> torch.Tensor:
|
||||
forwarded_states = self.layer_norm(hidden_states)
|
||||
forwarded_states = self.DenseReluDense(forwarded_states)
|
||||
hidden_states = hidden_states + forwarded_states
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5Attention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: T5Config,
|
||||
attn_type: AttentionType,
|
||||
has_relative_attention_bias=False,
|
||||
cache_config: Optional[CacheConfig] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.attn_type = attn_type
|
||||
# Cross-attention has no relative pos encoding anyway
|
||||
self.is_decoder = attn_type == AttentionType.DECODER
|
||||
self.has_relative_attention_bias = has_relative_attention_bias
|
||||
self.relative_attention_num_buckets = \
|
||||
config.relative_attention_num_buckets
|
||||
self.relative_attention_max_distance = \
|
||||
config.relative_attention_max_distance
|
||||
self.d_model = config.d_model
|
||||
self.key_value_proj_dim = config.d_kv
|
||||
assert cache_config
|
||||
# Alternatively we can get it from kv_cache size in fwd.
|
||||
self.block_size = cache_config.block_size
|
||||
|
||||
# Partition heads across multiple tensor parallel GPUs.
|
||||
tp_world_size = get_tensor_model_parallel_world_size()
|
||||
assert config.num_heads % tp_world_size == 0
|
||||
self.n_heads = config.num_heads // tp_world_size
|
||||
|
||||
self.inner_dim = self.n_heads * self.key_value_proj_dim
|
||||
# No GQA in t5.
|
||||
self.n_kv_heads = self.n_heads
|
||||
|
||||
self.qkv_proj = QKVParallelLinear(self.d_model,
|
||||
self.d_model // self.n_heads,
|
||||
self.n_heads,
|
||||
self.n_kv_heads,
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
|
||||
# NOTE (NickLucche) T5 employs a scaled weight initialization scheme
|
||||
# instead of scaling attention scores directly.
|
||||
self.attn = Attention(self.n_heads,
|
||||
config.d_kv,
|
||||
1.0,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.attn",
|
||||
attn_type=self.attn_type)
|
||||
|
||||
# Only the first SelfAttention block in encoder decoder has this
|
||||
# embedding layer, the others reuse its output.
|
||||
if self.has_relative_attention_bias:
|
||||
self.relative_attention_bias = \
|
||||
VocabParallelEmbedding(self.relative_attention_num_buckets,
|
||||
self.n_heads,
|
||||
org_num_embeddings=\
|
||||
self.relative_attention_num_buckets,
|
||||
quant_config=quant_config)
|
||||
self.out_proj = RowParallelLinear(
|
||||
self.inner_dim,
|
||||
self.d_model,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _relative_position_bucket(relative_position,
|
||||
bidirectional=True,
|
||||
num_buckets=32,
|
||||
max_distance=128):
|
||||
"""
|
||||
Adapted from Mesh Tensorflow:
|
||||
https://github.com/tensorflow/mesh/blob/0cb87fe07da627bf0b7e60475d59f95ed6b5be3d/mesh_tensorflow/transformer/transformer_layers.py#L593
|
||||
Translate relative position to a bucket number for relative attention.
|
||||
The relative position is defined as memory_position - query_position,
|
||||
i.e. the distance in tokens from the attending position to the
|
||||
attended-to position. If bidirectional=False, then positive relative
|
||||
positions are invalid. We use smaller buckets for small absolute
|
||||
relative_position and larger buckets for larger absolute
|
||||
relative_positions. All relative positions >=max_distance map to the
|
||||
same bucket. All relative positions <=-max_distance map to the same
|
||||
bucket. This should allow for more graceful generalization to longer
|
||||
sequences than the model has been trained on
|
||||
Args:
|
||||
relative_position: an int32 Tensor
|
||||
bidirectional: a boolean - whether the attention is bidirectional
|
||||
num_buckets: an integer
|
||||
max_distance: an integer
|
||||
Returns:
|
||||
a Tensor with the same shape as relative_position, containing int32
|
||||
values in the range [0, num_buckets)
|
||||
"""# noqa: E501
|
||||
relative_buckets = 0
|
||||
if bidirectional:
|
||||
num_buckets //= 2
|
||||
relative_buckets += (relative_position > 0).to(
|
||||
torch.long) * num_buckets
|
||||
relative_position = torch.abs(relative_position)
|
||||
else:
|
||||
relative_position = -torch.min(relative_position,
|
||||
torch.zeros_like(relative_position))
|
||||
# now relative_position is in the range [0, inf)
|
||||
|
||||
# half of the buckets are for exact increments in positions
|
||||
max_exact = num_buckets // 2
|
||||
is_small = relative_position < max_exact
|
||||
|
||||
# The other half of the buckets are for logarithmically bigger bins
|
||||
# in positions up to max_distance
|
||||
relative_position_if_large = max_exact + (
|
||||
torch.log(relative_position.float() / max_exact) /
|
||||
math.log(max_distance / max_exact) *
|
||||
(num_buckets - max_exact)).to(torch.long)
|
||||
relative_position_if_large = torch.min(
|
||||
relative_position_if_large,
|
||||
torch.full_like(relative_position_if_large, num_buckets - 1))
|
||||
|
||||
relative_buckets += torch.where(is_small, relative_position,
|
||||
relative_position_if_large)
|
||||
return relative_buckets
|
||||
|
||||
def compute_bias(self,
|
||||
query_length,
|
||||
key_length,
|
||||
device=None) -> torch.Tensor:
|
||||
"""Compute binned relative position bias"""
|
||||
# TODO possible tp issue?
|
||||
if device is None:
|
||||
device = self.relative_attention_bias.weight.device
|
||||
context_position = torch.arange(query_length,
|
||||
dtype=torch.long,
|
||||
device=device)[:, None]
|
||||
memory_position = torch.arange(key_length,
|
||||
dtype=torch.long,
|
||||
device=device)[None, :]
|
||||
# max_seq_len, nh
|
||||
relative_position = memory_position - context_position
|
||||
relative_position_bucket = self._relative_position_bucket(
|
||||
relative_position, # shape (query_length, key_length)
|
||||
bidirectional=(not self.is_decoder),
|
||||
num_buckets=self.relative_attention_num_buckets,
|
||||
max_distance=self.relative_attention_max_distance,
|
||||
)
|
||||
values = self.relative_attention_bias(
|
||||
relative_position_bucket
|
||||
) # shape (query_length, key_length, num_heads)
|
||||
x = values.permute([2, 0, 1]).unsqueeze(
|
||||
0) # shape (1, num_heads, query_length, key_length)
|
||||
return x
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor, # (num_tokens, d_model)
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
# TODO auto-selection of xformers backend when t5 is detected
|
||||
assert isinstance(attn_metadata, XFormersMetadata)
|
||||
num_seqs = len(
|
||||
attn_metadata.seq_lens) if attn_metadata.seq_lens else len(
|
||||
attn_metadata.encoder_seq_lens)
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
# Projection of 'own' hidden state (self-attention). No GQA here.
|
||||
q, k, v = qkv.split(self.inner_dim, dim=-1)
|
||||
|
||||
# NOTE (NickLucche) Attn bias is computed once per encoder or decoder
|
||||
# forward, on the first call to T5Attention.forward. Subsequent
|
||||
# *self-attention* layers will reuse it.
|
||||
attn_bias = _get_attn_bias(attn_metadata, self.attn_type)
|
||||
if self.attn_type == AttentionType.ENCODER_DECODER:
|
||||
# Projection of encoder's hidden states, cross-attention.
|
||||
if encoder_hidden_states is None:
|
||||
# Decode phase, kv already cached
|
||||
assert attn_metadata.num_prefills == 0
|
||||
k = None
|
||||
v = None
|
||||
else:
|
||||
assert attn_metadata.num_prefills > 0
|
||||
# Prefill phase (first decoder forward), caching kv
|
||||
qkv_enc, _ = self.qkv_proj(encoder_hidden_states)
|
||||
_, k, v = qkv_enc.split(self.inner_dim, dim=-1)
|
||||
# No custom attention bias must be set when running cross attn.
|
||||
assert attn_bias is None
|
||||
|
||||
# Not compatible with CP here (as all encoder-decoder models),
|
||||
# as it assumes homogeneous batch (prefills or decodes).
|
||||
elif self.has_relative_attention_bias:
|
||||
assert attn_bias is None # to be recomputed
|
||||
# Self-attention. Compute T5 relative positional encoding.
|
||||
# The bias term is computed on longest sequence in batch. Biases
|
||||
# for shorter sequences are slices of the longest.
|
||||
# TODO xformers-specific code.
|
||||
align_to = 8
|
||||
# bias expected shape: (num_seqs, NH, L, L_pad) for prefill,
|
||||
# (num_seqs, NH, 1, L_pad) for decodes.
|
||||
if self.attn_type == AttentionType.ENCODER:
|
||||
# Encoder prefill stage, uses xFormers, hence sequence
|
||||
# padding/alignment to 8 is required.
|
||||
seq_len = attn_metadata.max_encoder_seq_len
|
||||
padded_seq_len = (seq_len + align_to -
|
||||
1) // align_to * align_to
|
||||
# TODO (NickLucche) avoid extra copy on repeat,
|
||||
# provide multiple slices of same memory
|
||||
position_bias = self.compute_bias(seq_len,
|
||||
padded_seq_len).repeat(
|
||||
num_seqs, 1, 1, 1)
|
||||
# xFormers expects a list of biases, one matrix per sequence.
|
||||
# As each sequence gets its own bias, no masking is required.
|
||||
attn_bias = [
|
||||
p[None, :, :sq, :sq] for p, sq in zip(
|
||||
position_bias, attn_metadata.encoder_seq_lens)
|
||||
]
|
||||
elif attn_metadata.prefill_metadata:
|
||||
# Decoder prefill stage, uses xFormers, hence sequence
|
||||
# padding/alignment to 8 is required. First decoder step,
|
||||
# seq_len is usually 1, but one can prepend different start
|
||||
# tokens prior to generation.
|
||||
seq_len = attn_metadata.max_prefill_seq_len
|
||||
# ->align
|
||||
padded_seq_len = (seq_len + align_to -
|
||||
1) // align_to * align_to
|
||||
position_bias = self.compute_bias(seq_len,
|
||||
padded_seq_len).repeat(
|
||||
num_seqs, 1, 1, 1)
|
||||
# Causal mask for prefill.
|
||||
attn_bias = [
|
||||
LowerTriangularMaskWithTensorBias(pb[None, :, :sq, :sq])
|
||||
for pb, sq in zip(position_bias, attn_metadata.seq_lens)
|
||||
]
|
||||
else:
|
||||
# Decoder decoding stage, uses PagedAttention, hence sequence
|
||||
# padding/alignment to `block_size` is required. Expected
|
||||
# number of queries is always 1 (MQA not supported).
|
||||
seq_len = attn_metadata.max_decode_seq_len
|
||||
block_aligned_seq_len = (seq_len + self.block_size - 1
|
||||
) // self.block_size * self.block_size
|
||||
|
||||
# TODO bf16 bias support in PagedAttention.
|
||||
position_bias = self.compute_bias(
|
||||
seq_len, block_aligned_seq_len).float()
|
||||
# Bias for the last query, the one at current decoding step.
|
||||
position_bias = position_bias[:, :, -1:, :].repeat(
|
||||
num_seqs, 1, 1, 1)
|
||||
# No explicit masking required, this is done inside the
|
||||
# paged attention kernel based on the sequence length.
|
||||
attn_bias = [position_bias]
|
||||
|
||||
# NOTE Assign bias term on metadata based on attn type:
|
||||
# ENCODER->`encoder_attn_bias`, DECODER->`attn_bias`.
|
||||
_set_attn_bias(attn_metadata, attn_bias, self.attn_type)
|
||||
elif not self.has_relative_attention_bias:
|
||||
# Encoder/Decoder Self-Attention Layer, attn bias already cached.
|
||||
assert attn_bias is not None
|
||||
|
||||
attn_output = self.attn(q, k, v, kv_cache, attn_metadata)
|
||||
output, _ = self.out_proj(attn_output)
|
||||
return output
|
||||
|
||||
|
||||
class T5LayerSelfAttention(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
has_relative_attention_bias=False,
|
||||
cache_config: Optional[CacheConfig] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.SelfAttention = T5Attention(
|
||||
config,
|
||||
AttentionType.DECODER
|
||||
if "decoder" in prefix else AttentionType.ENCODER,
|
||||
has_relative_attention_bias=has_relative_attention_bias,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.SelfAttention")
|
||||
self.layer_norm = T5LayerNorm(config.d_model,
|
||||
eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
normed_hidden_states = self.layer_norm(hidden_states)
|
||||
attention_output = self.SelfAttention(
|
||||
hidden_states=normed_hidden_states,
|
||||
kv_cache=kv_cache,
|
||||
attn_metadata=attn_metadata,
|
||||
encoder_hidden_states=None,
|
||||
)
|
||||
hidden_states = hidden_states + attention_output
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5LayerCrossAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config,
|
||||
cache_config: Optional[CacheConfig] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.EncDecAttention = T5Attention(config,
|
||||
AttentionType.ENCODER_DECODER,
|
||||
has_relative_attention_bias=False,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.EncDecAttention")
|
||||
self.layer_norm = T5LayerNorm(config.d_model,
|
||||
eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
normed_hidden_states = self.layer_norm(hidden_states)
|
||||
attention_output = self.EncDecAttention(
|
||||
hidden_states=normed_hidden_states,
|
||||
kv_cache=kv_cache,
|
||||
attn_metadata=attn_metadata,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
)
|
||||
hidden_states = hidden_states + attention_output
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5Block(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: T5Config,
|
||||
is_decoder: bool,
|
||||
has_relative_attention_bias=False,
|
||||
cache_config: Optional[CacheConfig] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.is_decoder = is_decoder
|
||||
self.self_attn = T5LayerSelfAttention(
|
||||
config,
|
||||
has_relative_attention_bias=has_relative_attention_bias,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.self_attn")
|
||||
|
||||
if self.is_decoder:
|
||||
self.cross_attn = T5LayerCrossAttention(
|
||||
config,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.cross_attn")
|
||||
|
||||
self.ffn = T5LayerFF(config, quant_config=quant_config)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
|
||||
hidden_states = self.self_attn(
|
||||
hidden_states=hidden_states,
|
||||
kv_cache=kv_cache,
|
||||
attn_metadata=attn_metadata,
|
||||
)
|
||||
if self.is_decoder:
|
||||
hidden_states = self.cross_attn(
|
||||
hidden_states=hidden_states,
|
||||
kv_cache=kv_cache,
|
||||
attn_metadata=attn_metadata,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
)
|
||||
|
||||
# Apply Feed Forward layer
|
||||
hidden_states = self.ffn(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5Stack(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: T5Config,
|
||||
is_decoder: bool,
|
||||
n_layers: int,
|
||||
embed_tokens=None,
|
||||
cache_config: Optional[CacheConfig] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.embed_tokens = embed_tokens
|
||||
# Only the first block has relative positional encoding.
|
||||
self.blocks = nn.ModuleList([
|
||||
T5Block(config,
|
||||
is_decoder=is_decoder,
|
||||
has_relative_attention_bias=i == 0,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.blocks.{i}") for i in range(n_layers)
|
||||
])
|
||||
self.final_layer_norm = T5LayerNorm(config.d_model,
|
||||
eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
kv_caches: List[torch.Tensor],
|
||||
attn_metadata: AttentionMetadata,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None
|
||||
) -> torch.Tensor:
|
||||
hidden_states = self.embed_tokens(input_ids)
|
||||
|
||||
for idx, block in enumerate(self.blocks):
|
||||
hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
kv_cache=kv_caches[idx],
|
||||
attn_metadata=attn_metadata,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
)
|
||||
hidden_states = self.final_layer_norm(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5Model(nn.Module):
|
||||
_tied_weights_keys = [
|
||||
"encoder.embed_tokens.weight", "decoder.embed_tokens.weight"
|
||||
]
|
||||
|
||||
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
|
||||
super().__init__()
|
||||
config: T5Config = vllm_config.model_config.hf_config
|
||||
cache_config = vllm_config.cache_config
|
||||
quant_config = vllm_config.quant_config
|
||||
lora_config = vllm_config.lora_config
|
||||
|
||||
lora_vocab = (lora_config.lora_extra_vocab_size *
|
||||
(lora_config.max_loras or 1)) if lora_config else 0
|
||||
self.vocab_size = config.vocab_size + lora_vocab
|
||||
self.padding_idx = config.pad_token_id
|
||||
self.shared = VocabParallelEmbedding(
|
||||
config.vocab_size,
|
||||
config.d_model,
|
||||
org_num_embeddings=config.vocab_size)
|
||||
|
||||
self.encoder = T5Stack(config,
|
||||
False,
|
||||
config.num_layers,
|
||||
self.shared,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.encoder")
|
||||
self.decoder = T5Stack(config,
|
||||
True,
|
||||
config.num_decoder_layers,
|
||||
self.shared,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.decoder")
|
||||
|
||||
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
return self.shared(input_ids)
|
||||
|
||||
def forward(self, input_ids: torch.Tensor, encoder_input_ids: torch.Tensor,
|
||||
kv_caches: List[torch.Tensor],
|
||||
attn_metadata: AttentionMetadata) -> torch.Tensor:
|
||||
encoder_hidden_states = None
|
||||
|
||||
if encoder_input_ids.numel() > 0:
|
||||
# Run encoder attention if a non-zero number of encoder tokens
|
||||
# are provided as input: on a regular generate call, the encoder
|
||||
# runs once, on the prompt. Subsequent decoder calls reuse output
|
||||
# `encoder_hidden_states`.
|
||||
encoder_hidden_states = self.encoder(input_ids=encoder_input_ids,
|
||||
kv_caches=kv_caches,
|
||||
attn_metadata=attn_metadata)
|
||||
# Clear attention bias state.
|
||||
attn_metadata.attn_bias = None
|
||||
attn_metadata.encoder_attn_bias = None
|
||||
attn_metadata.cross_attn_bias = None
|
||||
|
||||
decoder_outputs = self.decoder(
|
||||
input_ids=input_ids,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
kv_caches=kv_caches,
|
||||
attn_metadata=attn_metadata)
|
||||
|
||||
# When capturing CUDA Graph
|
||||
attn_metadata.attn_bias = None
|
||||
attn_metadata.encoder_attn_bias = None
|
||||
attn_metadata.cross_attn_bias = None
|
||||
return decoder_outputs
|
||||
|
||||
|
||||
class T5ForConditionalGeneration(nn.Module):
|
||||
_keys_to_ignore_on_load_unexpected = [
|
||||
"decoder.block.0.layer.1.EncDecAttention.relative_attention_bias.weight",
|
||||
]
|
||||
_tied_weights_keys = [
|
||||
"encoder.embed_tokens.weight", "decoder.embed_tokens.weight",
|
||||
"lm_head.weight"
|
||||
]
|
||||
|
||||
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
|
||||
super().__init__()
|
||||
config: T5Config = vllm_config.model_config.hf_config
|
||||
self.model_dim = config.d_model
|
||||
self.config = config
|
||||
self.unpadded_vocab_size = config.vocab_size
|
||||
if lora_config := vllm_config.lora_config:
|
||||
self.unpadded_vocab_size += lora_config.lora_extra_vocab_size
|
||||
|
||||
self.model = T5Model(vllm_config=vllm_config,
|
||||
prefix=maybe_prefix(prefix, "model"))
|
||||
# Although not in config, this is the default for hf models.
|
||||
if self.config.tie_word_embeddings:
|
||||
self.lm_head = self.model.shared
|
||||
# in transformers this is smt more explicit, as in (after load)
|
||||
# self.lm_head.weight = self.model.shared.weight
|
||||
else:
|
||||
self.lm_head = ParallelLMHead(self.unpadded_vocab_size,
|
||||
config.d_model,
|
||||
org_num_embeddings=config.vocab_size)
|
||||
|
||||
self.logits_processor = LogitsProcessor(self.unpadded_vocab_size,
|
||||
config.vocab_size)
|
||||
self.sampler = get_sampler()
|
||||
|
||||
def compute_logits(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
sampling_metadata: SamplingMetadata,
|
||||
) -> Optional[torch.Tensor]:
|
||||
if self.config.tie_word_embeddings:
|
||||
# Rescale output before projecting on vocab
|
||||
# See https://github.com/tensorflow/mesh/blob/fa19d69eafc9a482aff0b59ddd96b025c0cb207d/mesh_tensorflow/transformer/transformer.py#L586 # noqa: E501
|
||||
hidden_states = hidden_states * (self.model_dim**-0.5)
|
||||
logits = self.logits_processor(self.lm_head, hidden_states,
|
||||
sampling_metadata)
|
||||
return logits
|
||||
|
||||
def sample(
|
||||
self,
|
||||
logits: Optional[torch.Tensor],
|
||||
sampling_metadata: SamplingMetadata,
|
||||
) -> Optional[SamplerOutput]:
|
||||
next_tokens = self.sampler(logits, sampling_metadata)
|
||||
return next_tokens
|
||||
|
||||
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
return self.model.shared(input_ids)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
kv_caches: List[torch.Tensor],
|
||||
attn_metadata: AttentionMetadata,
|
||||
intermediate_tensors: Optional[IntermediateTensors] = None,
|
||||
*,
|
||||
encoder_input_ids: torch.Tensor,
|
||||
encoder_positions: torch.Tensor,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
return self.model(input_ids, encoder_input_ids, kv_caches,
|
||||
attn_metadata)
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
||||
model_params_dict = dict(self.named_parameters(remove_duplicate=False))
|
||||
loaded_params: Set[str] = set()
|
||||
renamed_reg = [
|
||||
(re.compile(r'block\.(\d+)\.layer\.0'), r'blocks.\1.self_attn'),
|
||||
(re.compile(r'decoder.block\.(\d+)\.layer\.1'),
|
||||
r'decoder.blocks.\1.cross_attn'),
|
||||
(re.compile(r'decoder.block\.(\d+)\.layer\.2'),
|
||||
r'decoder.blocks.\1.ffn'),
|
||||
# encoder has no cross-attn, but rather self-attention+ffn.
|
||||
(re.compile(r'encoder.block\.(\d+)\.layer\.1'),
|
||||
r'encoder.blocks.\1.ffn'),
|
||||
(re.compile(r'\.o\.'), r'.out_proj.'),
|
||||
]
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj.", ".q.", "q"),
|
||||
(".qkv_proj.", ".k.", "k"),
|
||||
(".qkv_proj.", ".v.", "v")
|
||||
]
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
# No relative position attn bias on cross attention.
|
||||
if name in self._keys_to_ignore_on_load_unexpected:
|
||||
continue
|
||||
|
||||
# Handle some renaming
|
||||
for reg in renamed_reg:
|
||||
name = re.sub(*reg, name)
|
||||
|
||||
top_module, _ = name.split('.', 1)
|
||||
if top_module != 'lm_head':
|
||||
name = f"model.{name}"
|
||||
|
||||
# Split q/k/v layers to unified QKVParallelLinear
|
||||
for (param_name, weight_name, shard_id) in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
param = model_params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
break
|
||||
else:
|
||||
# Not a q/k/v layer.
|
||||
param = model_params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader",
|
||||
default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
return loaded_params
|
||||
@@ -0,0 +1,22 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from typing import List
|
||||
|
||||
def extract_layer_index(layer_name: str) -> int:
|
||||
"""
|
||||
Extract the layer index from the module name.
|
||||
Examples:
|
||||
- "encoder.layers.0" -> 0
|
||||
- "encoder.layers.1.self_attn" -> 1
|
||||
- "2.self_attn" -> 2
|
||||
- "model.encoder.layers.0.sub.1" -> ValueError
|
||||
"""
|
||||
subnames = layer_name.split(".")
|
||||
int_vals: List[int] = []
|
||||
for subname in subnames:
|
||||
try:
|
||||
int_vals.append(int(subname))
|
||||
except ValueError:
|
||||
continue
|
||||
assert len(int_vals) == 1, (f"layer name {layer_name} should"
|
||||
" only contain one integer")
|
||||
return int_vals[0]
|
||||
@@ -0,0 +1,150 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Final, Generic, Optional, Protocol, TypeVar, Union
|
||||
|
||||
import torch
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from vllm.attention.selector import (backend_name_to_enum,
|
||||
get_global_forced_attn_backend)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import _Backend, current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_C = TypeVar("_C", bound=PretrainedConfig)
|
||||
|
||||
|
||||
class VisionEncoderInfo(ABC, Generic[_C]):
|
||||
|
||||
def __init__(self, vision_config: _C) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.vision_config = vision_config
|
||||
|
||||
@abstractmethod
|
||||
def get_num_image_tokens(
|
||||
self,
|
||||
*,
|
||||
image_width: int,
|
||||
image_height: int,
|
||||
) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def get_max_image_tokens(self) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def get_image_size(self) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def get_patch_size(self) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def get_patch_grid_length(self) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class VisionLanguageConfig(Protocol):
|
||||
vision_config: Final[PretrainedConfig]
|
||||
|
||||
|
||||
def get_vision_encoder_info(
|
||||
hf_config: VisionLanguageConfig) -> VisionEncoderInfo:
|
||||
# Avoid circular imports
|
||||
from .clip import CLIPEncoderInfo, CLIPVisionConfig
|
||||
from .pixtral import PixtralHFEncoderInfo, PixtralVisionConfig
|
||||
from .siglip import SiglipEncoderInfo, SiglipVisionConfig
|
||||
|
||||
vision_config = hf_config.vision_config
|
||||
if isinstance(vision_config, CLIPVisionConfig):
|
||||
return CLIPEncoderInfo(vision_config)
|
||||
if isinstance(vision_config, PixtralVisionConfig):
|
||||
return PixtralHFEncoderInfo(vision_config)
|
||||
if isinstance(vision_config, SiglipVisionConfig):
|
||||
return SiglipEncoderInfo(vision_config)
|
||||
|
||||
msg = f"Unsupported vision config: {type(vision_config)}"
|
||||
raise NotImplementedError(msg)
|
||||
|
||||
|
||||
def get_vit_attn_backend(support_fa: bool = False) -> _Backend:
|
||||
"""
|
||||
Get the available attention backend for Vision Transformer.
|
||||
"""
|
||||
# TODO(Isotr0py): Remove `support_fa` after support FA for all ViTs attn.
|
||||
selected_backend: Optional[_Backend] = get_global_forced_attn_backend()
|
||||
if selected_backend is None:
|
||||
backend_by_env_var: Optional[str] = envs.VLLM_ATTENTION_BACKEND
|
||||
if backend_by_env_var is not None:
|
||||
selected_backend = backend_name_to_enum(backend_by_env_var)
|
||||
if selected_backend is None:
|
||||
if current_platform.is_cuda():
|
||||
device_available = current_platform.has_device_capability(80)
|
||||
if device_available and support_fa:
|
||||
from transformers.utils import is_flash_attn_2_available
|
||||
if is_flash_attn_2_available():
|
||||
selected_backend = _Backend.FLASH_ATTN
|
||||
else:
|
||||
logger.warning_once(
|
||||
"Current `vllm-flash-attn` has a bug inside vision "
|
||||
"module, so we use xformers backend instead. You can "
|
||||
"run `pip install flash-attn` to use flash-attention "
|
||||
"backend.")
|
||||
selected_backend = _Backend.XFORMERS
|
||||
else:
|
||||
# For Volta and Turing GPUs, use xformers instead.
|
||||
selected_backend = _Backend.XFORMERS
|
||||
else:
|
||||
# Default to torch SDPA for other non-GPU platforms.
|
||||
selected_backend = _Backend.TORCH_SDPA
|
||||
return selected_backend
|
||||
|
||||
|
||||
def resolve_visual_encoder_outputs(
|
||||
encoder_outputs: Union[torch.Tensor, list[torch.Tensor]],
|
||||
feature_sample_layers: Optional[list[int]],
|
||||
post_layer_norm: Optional[torch.nn.LayerNorm],
|
||||
max_possible_layers: int,
|
||||
) -> torch.Tensor:
|
||||
"""Given the outputs a visual encoder module that may correspond to the
|
||||
output of the last layer, or a list of hidden states to be stacked,
|
||||
handle post normalization and resolve it into a single output tensor.
|
||||
|
||||
Args:
|
||||
encoder_outputs: Output of encoder's last layer or all hidden states.
|
||||
feature_sample_layers: Optional layer indices to grab from the encoder
|
||||
outputs; if provided, encoder outputs must be a list.
|
||||
post_layer_norm: Post norm to apply to the output of the encoder.
|
||||
max_possible_layers: Total layers in the fully loaded visual encoder.
|
||||
|
||||
"""
|
||||
if feature_sample_layers is None:
|
||||
if post_layer_norm is not None:
|
||||
return post_layer_norm(encoder_outputs)
|
||||
return encoder_outputs
|
||||
|
||||
# Get the hidden states corresponding to the layer indices.
|
||||
# Negative values are relative to the full visual encoder,
|
||||
# so offset them depending on how many layers were loaded.
|
||||
# NOTE: this assumes that encoder_outputs is a list containing
|
||||
# the inputs to the visual encoder, followed by the hidden states
|
||||
# of each layer.
|
||||
num_loaded_layers = len(encoder_outputs) - 1
|
||||
offset = max_possible_layers - num_loaded_layers
|
||||
hs_pool = [
|
||||
encoder_outputs[layer_idx]
|
||||
if layer_idx >= 0 else encoder_outputs[layer_idx + offset]
|
||||
for layer_idx in feature_sample_layers
|
||||
]
|
||||
|
||||
# Apply post-norm on the final hidden state if we are using it
|
||||
uses_last_layer = feature_sample_layers[-1] in (len(hs_pool) - 1, -1)
|
||||
if post_layer_norm is not None and uses_last_layer:
|
||||
hs_pool[-1] = post_layer_norm(encoder_outputs)
|
||||
return torch.cat(hs_pool, dim=-1)
|
||||
@@ -0,0 +1,259 @@
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Utilities for Huggingface Transformers."""
|
||||
|
||||
import contextlib
|
||||
import os
|
||||
import warnings
|
||||
from pathlib import Path
|
||||
from typing import Dict, Optional, Type, Union, Any
|
||||
import json
|
||||
|
||||
from huggingface_hub import snapshot_download
|
||||
from transformers import (
|
||||
AutoConfig,
|
||||
AutoProcessor,
|
||||
AutoTokenizer,
|
||||
PretrainedConfig,
|
||||
PreTrainedTokenizer,
|
||||
PreTrainedTokenizerFast,
|
||||
)
|
||||
from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_NAMES
|
||||
|
||||
# from fastvideo.v1.models.configs import ChatGLMConfig, DbrxConfig, ExaoneConfig, Qwen2_5_VLConfig
|
||||
|
||||
_CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
|
||||
# ChatGLMConfig.model_type: ChatGLMConfig,
|
||||
# DbrxConfig.model_type: DbrxConfig,
|
||||
# ExaoneConfig.model_type: ExaoneConfig,
|
||||
# Qwen2_5_VLConfig.model_type: Qwen2_5_VLConfig,
|
||||
}
|
||||
|
||||
for name, cls in _CONFIG_REGISTRY.items():
|
||||
with contextlib.suppress(ValueError):
|
||||
AutoConfig.register(name, cls)
|
||||
|
||||
|
||||
def download_from_hf(model_path: str):
|
||||
if os.path.exists(model_path):
|
||||
return model_path
|
||||
|
||||
return snapshot_download(model_path, allow_patterns=["*.json", "*.bin", "*.model"])
|
||||
|
||||
|
||||
def get_hf_config(
|
||||
model: str,
|
||||
trust_remote_code: bool,
|
||||
revision: Optional[str] = None,
|
||||
model_override_args: Optional[dict] = None,
|
||||
inference_args: Optional[dict] = None,
|
||||
**kwargs,
|
||||
):
|
||||
is_gguf = check_gguf_file(model)
|
||||
if is_gguf:
|
||||
kwargs["gguf_file"] = model
|
||||
model = Path(model).parent
|
||||
|
||||
config = AutoConfig.from_pretrained(
|
||||
model, trust_remote_code=trust_remote_code, revision=revision, **kwargs
|
||||
)
|
||||
if config.model_type in _CONFIG_REGISTRY:
|
||||
config_class = _CONFIG_REGISTRY[config.model_type]
|
||||
config = config_class.from_pretrained(model, revision=revision)
|
||||
# NOTE(HandH1998): Qwen2VL requires `_name_or_path` attribute in `config`.
|
||||
setattr(config, "_name_or_path", model)
|
||||
if model_override_args:
|
||||
config.update(model_override_args)
|
||||
|
||||
# Special architecture mapping check for GGUF models
|
||||
if is_gguf:
|
||||
if config.model_type not in MODEL_FOR_CAUSAL_LM_MAPPING_NAMES:
|
||||
raise RuntimeError(f"Can't get gguf config for {config.model_type}.")
|
||||
model_type = MODEL_FOR_CAUSAL_LM_MAPPING_NAMES[config.model_type]
|
||||
config.update({"architectures": [model_type]})
|
||||
|
||||
return config
|
||||
|
||||
|
||||
def get_diffusers_config(
|
||||
model: str,
|
||||
inference_args: Optional[dict] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Gets a configuration for the given diffusers model.
|
||||
|
||||
Args:
|
||||
model: The model name or path.
|
||||
inference_args: Optional inference arguments to override in the config.
|
||||
|
||||
Returns:
|
||||
The loaded configuration.
|
||||
"""
|
||||
# Check if the model path exists
|
||||
if os.path.exists(model):
|
||||
config_file = os.path.join(model, "config.json")
|
||||
if os.path.exists(config_file):
|
||||
try:
|
||||
# Load the config directly from the file
|
||||
with open(config_file, "r") as f:
|
||||
config_dict = json.load(f)
|
||||
|
||||
# TODO(will): apply any overrides from inference args
|
||||
return config_dict
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to load diffusers config from {config_file}: {e}")
|
||||
else:
|
||||
raise RuntimeError(f"Diffusers config file not found at {model}")
|
||||
|
||||
|
||||
# Models don't use the same configuration key for determining the maximum
|
||||
# context length. Store them here so we can sanely check them.
|
||||
# NOTE: The ordering here is important. Some models have two of these and we
|
||||
# have a preference for which value gets used.
|
||||
CONTEXT_LENGTH_KEYS = [
|
||||
"max_sequence_length",
|
||||
"seq_length",
|
||||
"max_seq_len",
|
||||
"model_max_length",
|
||||
"max_position_embeddings",
|
||||
]
|
||||
|
||||
|
||||
def get_context_length(config):
|
||||
"""Get the context length of a model from a huggingface model configs."""
|
||||
text_config = config
|
||||
rope_scaling = getattr(text_config, "rope_scaling", None)
|
||||
if rope_scaling:
|
||||
rope_scaling_factor = rope_scaling.get("factor", 1)
|
||||
if "original_max_position_embeddings" in rope_scaling:
|
||||
rope_scaling_factor = 1
|
||||
if rope_scaling.get("rope_type", None) == "llama3":
|
||||
rope_scaling_factor = 1
|
||||
else:
|
||||
rope_scaling_factor = 1
|
||||
|
||||
for key in CONTEXT_LENGTH_KEYS:
|
||||
val = getattr(text_config, key, None)
|
||||
if val is not None:
|
||||
return int(rope_scaling_factor * val)
|
||||
return 2048
|
||||
|
||||
|
||||
# A fast LLaMA tokenizer with the pre-processed `tokenizer.json` file.
|
||||
_FAST_LLAMA_TOKENIZER = "hf-internal-testing/llama-tokenizer"
|
||||
|
||||
|
||||
def get_tokenizer(
|
||||
tokenizer_name: str,
|
||||
*args,
|
||||
tokenizer_mode: str = "auto",
|
||||
trust_remote_code: bool = False,
|
||||
tokenizer_revision: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast]:
|
||||
"""Gets a tokenizer for the given model name via Huggingface."""
|
||||
if tokenizer_mode == "slow":
|
||||
if kwargs.get("use_fast", False):
|
||||
raise ValueError("Cannot use the fast tokenizer in slow tokenizer mode.")
|
||||
kwargs["use_fast"] = False
|
||||
|
||||
is_gguf = check_gguf_file(tokenizer_name)
|
||||
if is_gguf:
|
||||
kwargs["gguf_file"] = tokenizer_name
|
||||
tokenizer_name = Path(tokenizer_name).parent
|
||||
|
||||
try:
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
tokenizer_name,
|
||||
*args,
|
||||
trust_remote_code=trust_remote_code,
|
||||
tokenizer_revision=tokenizer_revision,
|
||||
clean_up_tokenization_spaces=False,
|
||||
**kwargs,
|
||||
)
|
||||
except TypeError as e:
|
||||
# The LLaMA tokenizer causes a protobuf error in some environments.
|
||||
err_msg = (
|
||||
"Failed to load the tokenizer. If you are using a LLaMA V1 model "
|
||||
f"consider using '{_FAST_LLAMA_TOKENIZER}' instead of the "
|
||||
"original tokenizer."
|
||||
)
|
||||
raise RuntimeError(err_msg) from e
|
||||
except ValueError as e:
|
||||
# If the error pertains to the tokenizer class not existing or not
|
||||
# currently being imported, suggest using the --trust-remote-code flag.
|
||||
if not trust_remote_code and (
|
||||
"does not exist or is not currently imported." in str(e)
|
||||
or "requires you to execute the tokenizer file" in str(e)
|
||||
):
|
||||
err_msg = (
|
||||
"Failed to load the tokenizer. If the tokenizer is a custom "
|
||||
"tokenizer not yet available in the HuggingFace transformers "
|
||||
"library, consider setting `trust_remote_code=True` in LLM "
|
||||
"or using the `--trust-remote-code` flag in the CLI."
|
||||
)
|
||||
raise RuntimeError(err_msg) from e
|
||||
else:
|
||||
raise e
|
||||
|
||||
if not isinstance(tokenizer, PreTrainedTokenizerFast):
|
||||
warnings.warn(
|
||||
"Using a slow tokenizer. This might cause a significant "
|
||||
"slowdown. Consider using a fast tokenizer instead."
|
||||
)
|
||||
|
||||
attach_additional_stop_token_ids(tokenizer)
|
||||
return tokenizer
|
||||
|
||||
|
||||
def get_processor(
|
||||
tokenizer_name: str,
|
||||
*args,
|
||||
tokenizer_mode: str = "auto",
|
||||
trust_remote_code: bool = False,
|
||||
tokenizer_revision: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
processor = AutoProcessor.from_pretrained(
|
||||
tokenizer_name,
|
||||
*args,
|
||||
trust_remote_code=trust_remote_code,
|
||||
tokenizer_revision=tokenizer_revision,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
attach_additional_stop_token_ids(processor.tokenizer)
|
||||
return processor
|
||||
|
||||
|
||||
def attach_additional_stop_token_ids(tokenizer):
|
||||
# Special handling for stop token <|eom_id|> generated by llama 3 tool use.
|
||||
if "<|eom_id|>" in tokenizer.get_added_vocab():
|
||||
tokenizer.additional_stop_token_ids = set(
|
||||
[tokenizer.get_added_vocab()["<|eom_id|>"]]
|
||||
)
|
||||
else:
|
||||
tokenizer.additional_stop_token_ids = None
|
||||
|
||||
|
||||
def check_gguf_file(model: Union[str, os.PathLike]) -> bool:
|
||||
"""Check if the file is a GGUF model."""
|
||||
model = Path(model)
|
||||
if not model.is_file():
|
||||
return False
|
||||
elif model.suffix == ".gguf":
|
||||
return True
|
||||
|
||||
with open(model, "rb") as f:
|
||||
header = f.read(4)
|
||||
return header == b"GGUF"
|
||||
@@ -0,0 +1,409 @@
|
||||
from abc import ABC, abstractmethod
|
||||
import dataclasses
|
||||
|
||||
import torch
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
import os
|
||||
import glob
|
||||
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
|
||||
from transformers import PretrainedConfig, AutoTokenizer
|
||||
from fastvideo.v1.models.hf_transformer_utils import get_hf_config, get_diffusers_config
|
||||
from fastvideo.v1.models import get_scheduler
|
||||
from fastvideo.v1.models.registry import ModelRegistry
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from typing import Tuple, List, Optional, Any, Generator
|
||||
import time
|
||||
import torch.nn as nn
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
from fastvideo.v1.models.loader.weight_utils import (
|
||||
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
|
||||
pt_weights_iterator,
|
||||
safetensors_weights_iterator)
|
||||
from fastvideo.v1.models.loader.utils import set_default_torch_dtype
|
||||
from typing import (Any, Dict, Generator, Iterable, List, Optional,
|
||||
Tuple, cast)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ComponentLoader(ABC):
|
||||
"""Base class for loading a specific type of model component."""
|
||||
|
||||
def __init__(self, device=None):
|
||||
self.device = device
|
||||
|
||||
@abstractmethod
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""
|
||||
Load the component based on the model path, architecture, and inference args.
|
||||
|
||||
Args:
|
||||
model_path: Path to the component model
|
||||
architecture: Architecture of the component model
|
||||
inference_args: Inference arguments
|
||||
|
||||
Returns:
|
||||
The loaded component
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def for_module_type(cls, module_type: str, transformers_or_diffusers: str) -> 'ComponentLoader':
|
||||
"""
|
||||
Factory method to create a component loader for a specific module type.
|
||||
|
||||
Args:
|
||||
module_type: Type of module (e.g., "vae", "text_encoder", "transformer", "scheduler")
|
||||
transformers_or_diffusers: Whether the module is from transformers or diffusers
|
||||
|
||||
Returns:
|
||||
A component loader for the specified module type
|
||||
"""
|
||||
# Map of module types to their loader classes and expected library
|
||||
module_loaders = {
|
||||
"scheduler": (SchedulerLoader, "diffusers"),
|
||||
"transformer": (TransformerLoader, "diffusers"),
|
||||
"vae": (VAELoader, "diffusers"),
|
||||
"text_encoder": (TextEncoderLoader, "transformers"),
|
||||
"text_encoder_2": (TextEncoderLoader, "transformers"),
|
||||
"tokenizer": (TokenizerLoader, "transformers"),
|
||||
"tokenizer_2": (TokenizerLoader, "transformers"),
|
||||
}
|
||||
|
||||
if module_type in module_loaders:
|
||||
loader_cls, expected_library = module_loaders[module_type]
|
||||
# Assert that the library matches what's expected for this module type
|
||||
assert transformers_or_diffusers == expected_library, f"{module_type} must be loaded from {expected_library}, got {transformers_or_diffusers}"
|
||||
return loader_cls()
|
||||
|
||||
# For unknown module types, use a generic loader
|
||||
logger.warning(f"No specific loader found for module type: {module_type}. Using generic loader.")
|
||||
return GenericComponentLoader(transformers_or_diffusers)
|
||||
|
||||
|
||||
class TextEncoderLoader(ComponentLoader):
|
||||
"""Loader for text encoders."""
|
||||
@dataclasses.dataclass
|
||||
class Source:
|
||||
"""A source for weights."""
|
||||
|
||||
model_or_path: str
|
||||
"""The model ID or path."""
|
||||
|
||||
prefix: str = ""
|
||||
"""A prefix to prepend to all weights."""
|
||||
|
||||
fall_back_to_pt: bool = True
|
||||
"""Whether .pt weights can be used."""
|
||||
|
||||
allow_patterns_overrides: Optional[list[str]] = None
|
||||
"""If defined, weights will load exclusively using these patterns."""
|
||||
|
||||
counter_before_loading_weights: float = 0.0
|
||||
counter_after_loading_weights: float = 0.0
|
||||
|
||||
def _prepare_weights(
|
||||
self,
|
||||
model_name_or_path: str,
|
||||
fall_back_to_pt: bool,
|
||||
allow_patterns_overrides: Optional[list[str]],
|
||||
) -> Tuple[str, List[str], bool]:
|
||||
"""Prepare weights for the model.
|
||||
|
||||
If the model is not local, it will be downloaded."""
|
||||
# model_name_or_path = (self._maybe_download_from_modelscope(
|
||||
# model_name_or_path, revision) or model_name_or_path)
|
||||
|
||||
is_local = os.path.isdir(model_name_or_path)
|
||||
assert is_local, "Model path must be a local directory"
|
||||
|
||||
use_safetensors = False
|
||||
index_file = SAFE_WEIGHTS_INDEX_NAME
|
||||
allow_patterns = ["*.safetensors", "*.bin"]
|
||||
|
||||
|
||||
if fall_back_to_pt:
|
||||
allow_patterns += ["*.pt"]
|
||||
|
||||
if allow_patterns_overrides is not None:
|
||||
allow_patterns = allow_patterns_overrides
|
||||
|
||||
|
||||
hf_folder = model_name_or_path
|
||||
|
||||
hf_weights_files: List[str] = []
|
||||
for pattern in allow_patterns:
|
||||
hf_weights_files += glob.glob(os.path.join(hf_folder, pattern))
|
||||
if len(hf_weights_files) > 0:
|
||||
if pattern == "*.safetensors":
|
||||
use_safetensors = True
|
||||
break
|
||||
|
||||
if use_safetensors:
|
||||
hf_weights_files = filter_duplicate_safetensors_files(
|
||||
hf_weights_files, hf_folder, index_file)
|
||||
else:
|
||||
hf_weights_files = filter_files_not_needed_for_inference(
|
||||
hf_weights_files)
|
||||
|
||||
if len(hf_weights_files) == 0:
|
||||
raise RuntimeError(
|
||||
f"Cannot find any model weights with `{model_name_or_path}`")
|
||||
|
||||
return hf_folder, hf_weights_files, use_safetensors
|
||||
|
||||
def _get_weights_iterator(
|
||||
self, source: "Source"
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Get an iterator for the model weights based on the load format."""
|
||||
hf_folder, hf_weights_files, use_safetensors = self._prepare_weights(
|
||||
source.model_or_path, source.fall_back_to_pt,
|
||||
source.allow_patterns_overrides)
|
||||
if use_safetensors:
|
||||
weights_iterator = safetensors_weights_iterator(hf_weights_files)
|
||||
else:
|
||||
weights_iterator = pt_weights_iterator(hf_weights_files)
|
||||
|
||||
|
||||
if self.counter_before_loading_weights == 0.0:
|
||||
self.counter_before_loading_weights = time.perf_counter()
|
||||
# Apply the prefix.
|
||||
return ((source.prefix + name, tensor)
|
||||
for (name, tensor) in weights_iterator)
|
||||
|
||||
def _get_all_weights(
|
||||
self,
|
||||
model_config: Dict[str, Any],
|
||||
model: nn.Module,
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
primary_weights = TextEncoderLoader.Source(
|
||||
model_config.model,
|
||||
prefix="",
|
||||
fall_back_to_pt=getattr(model, "fall_back_to_pt_during_load",
|
||||
True),
|
||||
allow_patterns_overrides=getattr(model, "allow_patterns_overrides",
|
||||
None),
|
||||
)
|
||||
yield from self._get_weights_iterator(primary_weights)
|
||||
|
||||
secondary_weights = cast(
|
||||
Iterable[TextEncoderLoader.Source],
|
||||
getattr(model, "secondary_weights", ()),
|
||||
)
|
||||
for source in secondary_weights:
|
||||
yield from self._get_weights_iterator(source)
|
||||
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load the text encoders based on the model path, architecture, and inference args."""
|
||||
model_config: PretrainedConfig = get_hf_config(
|
||||
model=model_path,
|
||||
trust_remote_code=inference_args.trust_remote_code,
|
||||
revision=inference_args.revision,
|
||||
model_override_args=None,
|
||||
inference_args=inference_args,
|
||||
)
|
||||
logger.info(f"HF Model config: {model_config}")
|
||||
|
||||
|
||||
target_device = torch.device(inference_args.device_str)
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, model_config, target_device)
|
||||
|
||||
def load_model(self, model_path: str, model_config, target_device: torch.device):
|
||||
with set_default_torch_dtype(torch.float16):
|
||||
with target_device:
|
||||
architectures = getattr(model_config, "architectures", [])
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
|
||||
model = model_cls(model_config)
|
||||
|
||||
weights_to_load = {name for name, _ in model.named_parameters()}
|
||||
model_config.model = model_path
|
||||
loaded_weights = model.load_weights(
|
||||
self._get_all_weights(model_config, model))
|
||||
self.counter_after_loading_weights = time.perf_counter()
|
||||
logger.info(
|
||||
"Loading weights took %.2f seconds",
|
||||
self.counter_after_loading_weights -
|
||||
self.counter_before_loading_weights)
|
||||
# We only enable strict check for non-quantized models
|
||||
# that have loaded weights tracking currently.
|
||||
# if loaded_weights is not None:
|
||||
weights_not_loaded = weights_to_load - loaded_weights
|
||||
if weights_not_loaded:
|
||||
raise ValueError(
|
||||
"Following weights were not initialized from "
|
||||
f"checkpoint: {weights_not_loaded}")
|
||||
|
||||
# TODO(will): add support for training/finetune
|
||||
return model.eval()
|
||||
|
||||
|
||||
class TokenizerLoader(ComponentLoader):
|
||||
"""Loader for tokenizers."""
|
||||
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load the tokenizer based on the model path, architecture, and inference args."""
|
||||
logger.info(f"Loading tokenizer from {model_path}")
|
||||
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
model_path,
|
||||
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
|
||||
# other method of config?
|
||||
padding_size='right',
|
||||
)
|
||||
logger.info(f"Loaded tokenizer: {tokenizer.__class__.__name__}")
|
||||
return tokenizer
|
||||
|
||||
|
||||
class VAELoader(ComponentLoader):
|
||||
"""Loader for VAE."""
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load the VAE based on the model path, architecture, and inference args."""
|
||||
# TODO(will): move this to a constants file
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
config = get_diffusers_config(model=model_path)
|
||||
|
||||
class_name = config.pop("_class_name")
|
||||
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
||||
config.pop("_diffusers_version")
|
||||
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
vae = vae_cls(**config).to(inference_args.device)
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
|
||||
# TODO(PY)
|
||||
assert len(safetensors_list) == 1, f"Found {len(safetensors_list)} safetensors files in {d}"
|
||||
loaded = safetensors_load_file(safetensors_list[0])
|
||||
vae.load_state_dict(loaded)
|
||||
dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
|
||||
vae = vae.eval().to(dtype)
|
||||
|
||||
# TODO(will): should we define hunyuan vae config class?
|
||||
vae_kwargs = {
|
||||
"s_ratio": config["spatial_compression_ratio"],
|
||||
"t_ratio": config["temporal_compression_ratio"],
|
||||
}
|
||||
|
||||
vae.kwargs = vae_kwargs
|
||||
|
||||
return vae
|
||||
|
||||
class TransformerLoader(ComponentLoader):
|
||||
"""Loader for transformer."""
|
||||
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load the transformer based on the model path, architecture, and inference args."""
|
||||
model_config = get_diffusers_config(model=model_path)
|
||||
cls_name = model_config.pop("_class_name")
|
||||
if cls_name is None:
|
||||
raise ValueError(f"Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported.")
|
||||
model_config.pop("_diffusers_version")
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
|
||||
if not safetensors_list:
|
||||
raise ValueError(f"No safetensors files found in {model_path}")
|
||||
|
||||
logger.info(f"Loading model from {len(safetensors_list)} safetensors files in {model_path}")
|
||||
|
||||
# initialize_sequence_parallel_group(inference_args.sp_size)
|
||||
|
||||
# Load the model using FSDP loader
|
||||
logger.info(f"Loading model from {cls_name}")
|
||||
model = load_fsdp_model(
|
||||
model_cls=model_cls,
|
||||
init_params=model_config,
|
||||
weight_dir_list=safetensors_list,
|
||||
device=inference_args.device,
|
||||
cpu_offload=inference_args.use_cpu_offload
|
||||
)
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info(f"Loaded model with {total_params / 1e9:.2f}B parameters")
|
||||
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
class SchedulerLoader(ComponentLoader):
|
||||
"""Loader for scheduler."""
|
||||
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load the scheduler based on the model path, architecture, and inference args."""
|
||||
|
||||
scheduler = get_scheduler(
|
||||
module_path=model_path,
|
||||
architecture=architecture,
|
||||
inference_args=inference_args,
|
||||
)
|
||||
logger.info(f"Scheduler loaded: {scheduler}")
|
||||
return scheduler
|
||||
|
||||
class GenericComponentLoader(ComponentLoader):
|
||||
"""Generic loader for components that don't have a specific loader."""
|
||||
|
||||
def __init__(self, library="transformers"):
|
||||
super().__init__()
|
||||
self.library = library
|
||||
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load a generic component based on the model path, architecture, and inference args."""
|
||||
logger.warning(f"Using generic loader for {model_path} with library {self.library}")
|
||||
|
||||
if self.library == "transformers":
|
||||
from transformers import AutoModel
|
||||
|
||||
model = AutoModel.from_pretrained(
|
||||
model_path,
|
||||
trust_remote_code=inference_args.trust_remote_code,
|
||||
revision=inference_args.revision,
|
||||
)
|
||||
logger.info(f"Loaded generic transformers model: {model.__class__.__name__}")
|
||||
return model
|
||||
elif self.library == "diffusers":
|
||||
logger.warning(f"Generic loading for diffusers components is not fully implemented")
|
||||
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
|
||||
|
||||
model_config = get_diffusers_config(model=model_path)
|
||||
logger.info(f"Diffusers Model config: {model_config}")
|
||||
# This is a placeholder - in a real implementation, you'd need to handle this properly
|
||||
return None
|
||||
else:
|
||||
raise ValueError(f"Unsupported library: {self.library}")
|
||||
|
||||
class PipelineComponentLoader:
|
||||
"""
|
||||
Utility class for loading pipeline components.
|
||||
This replaces the chain of if-else statements in load_pipeline_module.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def load_module(module_name: str, component_model_path: str, transformers_or_diffusers: str,
|
||||
architecture: str, inference_args: InferenceArgs):
|
||||
"""
|
||||
Load a pipeline module.
|
||||
|
||||
Args:
|
||||
module_name: Name of the module (e.g., "vae", "text_encoder", "transformer", "scheduler")
|
||||
component_model_path: Path to the component model
|
||||
transformers_or_diffusers: Whether the module is from transformers or diffusers
|
||||
architecture: Architecture of the component model
|
||||
inference_args: Inference arguments
|
||||
|
||||
Returns:
|
||||
The loaded module
|
||||
"""
|
||||
logger.info(f"Loading {module_name} using {transformers_or_diffusers} from {component_model_path}")
|
||||
|
||||
# Get the appropriate loader for this module type
|
||||
loader = ComponentLoader.for_module_type(module_name, transformers_or_diffusers)
|
||||
|
||||
# Load the module
|
||||
return loader.load(component_model_path, architecture, inference_args)
|
||||
@@ -0,0 +1,227 @@
|
||||
# Adapted from torchtune
|
||||
# Copyright 2024 The TorchTune Authors.
|
||||
# Copyright 2025 The FastVideo Authors.
|
||||
|
||||
from collections import defaultdict
|
||||
from typing import Any, Callable, Dict, Generator, List, Optional, Tuple, Type
|
||||
from itertools import chain
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.distributed import DeviceMesh, init_device_mesh
|
||||
from fastvideo.v1.distributed.parallel_state import get_sequence_model_parallel_world_size
|
||||
from torch.distributed._composable.fsdp import CPUOffloadPolicy, fully_shard
|
||||
from torch.distributed._tensor import distribute_tensor
|
||||
from torch.nn.modules.module import _IncompatibleKeys
|
||||
from vllm.model_executor.model_loader.weight_utils import safetensors_weights_iterator
|
||||
|
||||
import contextlib
|
||||
import re
|
||||
|
||||
# TODO(PY): move this to utils elsewhere
|
||||
@contextlib.contextmanager
|
||||
def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
|
||||
"""
|
||||
Context manager to set torch's default dtype.
|
||||
|
||||
Args:
|
||||
dtype (torch.dtype): The desired default dtype inside the context manager.
|
||||
|
||||
Returns:
|
||||
ContextManager: context manager for setting default dtype.
|
||||
|
||||
Example:
|
||||
>>> with set_default_dtype(torch.bfloat16):
|
||||
>>> x = torch.tensor([1, 2, 3])
|
||||
>>> x.dtype
|
||||
torch.bfloat16
|
||||
|
||||
|
||||
"""
|
||||
old_dtype = torch.get_default_dtype()
|
||||
torch.set_default_dtype(dtype)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
torch.set_default_dtype(old_dtype)
|
||||
|
||||
|
||||
def get_param_names_mapping(mapping_dict: Dict[str, str]) -> Callable[[str], str]:
|
||||
"""
|
||||
Creates a mapping function that transforms parameter names using regex patterns.
|
||||
|
||||
Args:
|
||||
mapping_dict (Dict[str, str]): Dictionary mapping regex patterns to replacement patterns
|
||||
param_name (str): The parameter name to be transformed
|
||||
|
||||
Returns:
|
||||
Callable[[str], str]: A function that maps parameter names from source to target format
|
||||
"""
|
||||
def mapping_fn(name: str) -> str:
|
||||
|
||||
# Try to match and transform the name using the regex patterns in mapping_dict
|
||||
for pattern, replacement in mapping_dict.items():
|
||||
match = re.match(pattern, name)
|
||||
if match:
|
||||
merge_index = None
|
||||
total_splitted_params = None
|
||||
if isinstance(replacement, tuple):
|
||||
merge_index = replacement[1]
|
||||
total_splitted_params = replacement[2]
|
||||
replacement= replacement[0]
|
||||
name = re.sub(pattern, replacement, name)
|
||||
return name, merge_index, total_splitted_params
|
||||
|
||||
# If no pattern matches, return the original name
|
||||
return name, None, None
|
||||
|
||||
return mapping_fn
|
||||
|
||||
|
||||
# TODO(PY): add compile option
|
||||
def load_fsdp_model(
|
||||
model_cls: Type[nn.Module],
|
||||
init_params: Dict[str, Any],
|
||||
weight_dir_list: List[str],
|
||||
device: torch.device,
|
||||
cpu_offload: bool = False,
|
||||
default_dtype: Optional[torch.dtype] = torch.bfloat16,
|
||||
) -> torch.nn.Module:
|
||||
with set_default_dtype(default_dtype), torch.device("meta"):
|
||||
model = model_cls(**init_params)
|
||||
device_mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(get_sequence_model_parallel_world_size(),),
|
||||
mesh_dim_names=("dp", ),
|
||||
)
|
||||
shard_model(model, cpu_offload=cpu_offload, reshard_after_forward=True, dp_mesh=device_mesh["dp"])
|
||||
weight_iterator = safetensors_weights_iterator(weight_dir_list)
|
||||
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
|
||||
load_fsdp_model_from_full_model_state_dict(
|
||||
model,
|
||||
weight_iterator,
|
||||
device,
|
||||
strict=True,
|
||||
cpu_offload=cpu_offload,
|
||||
param_names_mapping=param_names_mapping_fn,
|
||||
)
|
||||
for n, p in chain(model.named_parameters(), model.named_buffers()):
|
||||
if p.is_meta:
|
||||
raise RuntimeError(f"Unexpected param or buffer {n} on meta device.")
|
||||
for p in model.parameters():
|
||||
p.requires_grad = False
|
||||
return model
|
||||
|
||||
def shard_model(
|
||||
model,
|
||||
*,
|
||||
cpu_offload: bool,
|
||||
reshard_after_forward: bool = True,
|
||||
dp_mesh: Optional[DeviceMesh] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API.
|
||||
|
||||
This method will over the model's named modules from the bottom-up and apply shard modules
|
||||
based on whether they meet any of the criteria from shard_conditions.
|
||||
|
||||
Args:
|
||||
model (TransformerDecoder): Model to shard with FSDP.
|
||||
shard_conditions (List[Callable[[str, nn.Module], bool]]): A list of functions to determine
|
||||
which modules to shard with FSDP. Each function should take module name (relative to root)
|
||||
and the module itself, returning True if FSDP should shard the module and False otherwise.
|
||||
If any of shard_conditions return True for a given module, it will be sharded by FSDP.
|
||||
cpu_offload (bool): If set to True, FSDP will offload parameters, gradients, and optimizer
|
||||
states to CPU.
|
||||
reshard_after_forward (bool): Whether to reshard parameters and buffers after
|
||||
the forward pass. Setting this to True corresponds to the FULL_SHARD sharding strategy
|
||||
from FSDP1, while setting it to False corresponds to the SHARD_GRAD_OP sharding strategy.
|
||||
dp_mesh (Optional[DeviceMesh]): Device mesh to use for FSDP sharding under mutliple parallelism.
|
||||
Default to None.
|
||||
|
||||
Raises:
|
||||
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
|
||||
"""
|
||||
fsdp_kwargs = {"reshard_after_forward": reshard_after_forward, "mesh": dp_mesh}
|
||||
if cpu_offload:
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
|
||||
|
||||
# Shard the model with FSDP, iterating in reverse to start with
|
||||
# lowest-level modules first
|
||||
num_layers_sharded = 0
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([shard_condition(n, m) for shard_condition in model._fsdp_shard_conditions]):
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
num_layers_sharded += 1
|
||||
|
||||
if num_layers_sharded == 0:
|
||||
raise ValueError(
|
||||
"No layer modules were sharded. Please check if shard conditions are working as expected."
|
||||
)
|
||||
|
||||
# Finally shard the entire model to account for any stragglers
|
||||
fully_shard(model, **fsdp_kwargs)
|
||||
|
||||
# TODO(PY): device mesh for cfg parallel
|
||||
def load_fsdp_model_from_full_model_state_dict(
|
||||
model: torch.nn.Module,
|
||||
full_sd_iterator: Generator[Tuple[str, torch.Tensor], None, None],
|
||||
device: torch.device,
|
||||
strict: bool = False,
|
||||
cpu_offload: bool = False,
|
||||
param_names_mapping: Optional[Callable[[str], str]] = None,
|
||||
) -> _IncompatibleKeys:
|
||||
"""
|
||||
Converting full state dict into a sharded state dict
|
||||
and loading it into FSDP model
|
||||
Args:
|
||||
model (FSDPModule): Model to generate fully qualified names for cpu_state_dict
|
||||
full_sd_iterator (Generator): an iterator yielding (param_name, tensor) pairs
|
||||
device (torch.device): device used to move full state dict tensors
|
||||
strict (bool): flag to check if to load the model in strict mode
|
||||
cpu_offload (bool): flag to check if offload to CPU is enabled
|
||||
param_names_mapping (Optional[Callable[[str], str]]): a function that maps full param name to sharded param name
|
||||
|
||||
Returns:
|
||||
``NamedTuple`` with ``missing_keys`` and ``unexpected_keys`` fields:
|
||||
* **missing_keys** is a list of str containing the missing keys
|
||||
* **unexpected_keys** is a list of str containing the unexpected keys
|
||||
|
||||
Raises:
|
||||
NotImplementedError: If got FSDP with more than 1D.
|
||||
"""
|
||||
meta_sharded_sd = model.state_dict()
|
||||
|
||||
sharded_sd = {}
|
||||
to_merge_params = defaultdict(dict)
|
||||
for source_param_name, full_tensor in full_sd_iterator:
|
||||
target_param_name, merge_index, num_params_to_merge = param_names_mapping(source_param_name)
|
||||
|
||||
if merge_index is not None:
|
||||
to_merge_params[target_param_name][merge_index] = full_tensor
|
||||
if len(to_merge_params[target_param_name]) == num_params_to_merge:
|
||||
# cat at dim=1 according to the merge_index order
|
||||
sorted_tensors = [to_merge_params[target_param_name][i] for i in range(num_params_to_merge)]
|
||||
full_tensor = torch.cat(sorted_tensors, dim=0)
|
||||
del to_merge_params[target_param_name]
|
||||
else:
|
||||
continue
|
||||
|
||||
sharded_meta_param = meta_sharded_sd.get(target_param_name)
|
||||
if sharded_meta_param is None:
|
||||
raise ValueError(f"Parameter {source_param_name}-->{target_param_name} not found in meta sharded state dict")
|
||||
full_tensor = full_tensor.to(sharded_meta_param.dtype).to(device)
|
||||
|
||||
if not hasattr(sharded_meta_param, "device_mesh"):
|
||||
# In cases where parts of the model aren't sharded, some parameters will be plain tensors
|
||||
sharded_tensor = full_tensor
|
||||
else:
|
||||
sharded_tensor = distribute_tensor(
|
||||
full_tensor,
|
||||
sharded_meta_param.device_mesh,
|
||||
sharded_meta_param.placements,
|
||||
)
|
||||
if cpu_offload:
|
||||
sharded_tensor = sharded_tensor.cpu()
|
||||
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
|
||||
# choose `assign=True` since we cannot call `copy_` on meta tensor
|
||||
return model.load_state_dict(sharded_sd, strict=strict, assign=True)
|
||||
@@ -0,0 +1,17 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Utilities for selecting and loading models."""
|
||||
import contextlib
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def set_default_torch_dtype(dtype: torch.dtype):
|
||||
"""Sets the default torch dtype to the given dtype."""
|
||||
old_dtype = torch.get_default_dtype()
|
||||
torch.set_default_dtype(dtype)
|
||||
yield
|
||||
torch.set_default_dtype(old_dtype)
|
||||
@@ -0,0 +1,346 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Adapted from vllm
|
||||
# Copyright 2023 The vLLM Authors.
|
||||
# Copyright 2025 The FastVideo Authors.
|
||||
|
||||
"""Utilities for downloading and initializing model weights."""
|
||||
import fnmatch
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import tempfile
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Generator, List, Optional, Tuple, Union
|
||||
|
||||
import filelock
|
||||
import huggingface_hub.constants
|
||||
import torch
|
||||
from huggingface_hub import HfFileSystem, hf_hub_download, snapshot_download
|
||||
from safetensors.torch import safe_open
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# use system-level temp directory for file locks, so that multiple users
|
||||
# can share the same lock without error.
|
||||
# lock files in the temp directory will be automatically deleted when the
|
||||
# system reboots, so users will not complain about annoying lock files
|
||||
temp_dir = tempfile.gettempdir()
|
||||
|
||||
|
||||
def enable_hf_transfer():
|
||||
"""automatically activates hf_transfer
|
||||
"""
|
||||
if "HF_HUB_ENABLE_HF_TRANSFER" not in os.environ:
|
||||
try:
|
||||
# enable hf hub transfer if available
|
||||
import hf_transfer # type: ignore # noqa
|
||||
huggingface_hub.constants.HF_HUB_ENABLE_HF_TRANSFER = True
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
|
||||
enable_hf_transfer()
|
||||
|
||||
|
||||
class DisabledTqdm(tqdm):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs, disable=True)
|
||||
|
||||
|
||||
def get_lock(model_name_or_path: Union[str, Path],
|
||||
cache_dir: Optional[str] = None):
|
||||
lock_dir = cache_dir or temp_dir
|
||||
model_name_or_path = str(model_name_or_path)
|
||||
os.makedirs(os.path.dirname(lock_dir), exist_ok=True)
|
||||
model_name = model_name_or_path.replace("/", "-")
|
||||
hash_name = hashlib.sha256(model_name.encode()).hexdigest()
|
||||
# add hash to avoid conflict with old users' lock files
|
||||
lock_file_name = hash_name + model_name + ".lock"
|
||||
# mode 0o666 is required for the filelock to be shared across users
|
||||
lock = filelock.FileLock(os.path.join(lock_dir, lock_file_name),
|
||||
mode=0o666)
|
||||
return lock
|
||||
|
||||
|
||||
def _shared_pointers(tensors):
|
||||
ptrs = defaultdict(list)
|
||||
for k, v in tensors.items():
|
||||
ptrs[v.data_ptr()].append(k)
|
||||
failing = []
|
||||
for _, names in ptrs.items():
|
||||
if len(names) > 1:
|
||||
failing.append(names)
|
||||
return failing
|
||||
|
||||
|
||||
def download_weights_from_hf(
|
||||
model_name_or_path: str,
|
||||
cache_dir: Optional[str],
|
||||
allow_patterns: List[str],
|
||||
revision: Optional[str] = None,
|
||||
ignore_patterns: Optional[Union[str, List[str]]] = None,
|
||||
) -> str:
|
||||
"""Download model weights from Hugging Face Hub.
|
||||
|
||||
Args:
|
||||
model_name_or_path (str): The model name or path.
|
||||
cache_dir (Optional[str]): The cache directory to store the model
|
||||
weights. If None, will use HF defaults.
|
||||
allow_patterns (List[str]): The allowed patterns for the
|
||||
weight files. Files matched by any of the patterns will be
|
||||
downloaded.
|
||||
revision (Optional[str]): The revision of the model.
|
||||
ignore_patterns (Optional[Union[str, List[str]]]): The patterns to
|
||||
filter out the weight files. Files matched by any of the patterns
|
||||
will be ignored.
|
||||
|
||||
Returns:
|
||||
str: The path to the downloaded model weights.
|
||||
"""
|
||||
local_only = huggingface_hub.constants.HF_HUB_OFFLINE
|
||||
if not local_only:
|
||||
# Before we download we look at that is available:
|
||||
fs = HfFileSystem()
|
||||
file_list = fs.ls(model_name_or_path, detail=False, revision=revision)
|
||||
|
||||
# depending on what is available we download different things
|
||||
for pattern in allow_patterns:
|
||||
matching = fnmatch.filter(file_list, pattern)
|
||||
if len(matching) > 0:
|
||||
allow_patterns = [pattern]
|
||||
break
|
||||
|
||||
logger.info("Using model weights format %s", allow_patterns)
|
||||
# Use file lock to prevent multiple processes from
|
||||
# downloading the same model weights at the same time.
|
||||
with get_lock(model_name_or_path, cache_dir):
|
||||
start_time = time.perf_counter()
|
||||
hf_folder = snapshot_download(
|
||||
model_name_or_path,
|
||||
allow_patterns=allow_patterns,
|
||||
ignore_patterns=ignore_patterns,
|
||||
cache_dir=cache_dir,
|
||||
tqdm_class=DisabledTqdm,
|
||||
revision=revision,
|
||||
local_files_only=local_only,
|
||||
)
|
||||
time_taken = time.perf_counter() - start_time
|
||||
if time_taken > 0.5:
|
||||
logger.info("Time spent downloading weights for %s: %.6f seconds",
|
||||
model_name_or_path, time_taken)
|
||||
return hf_folder
|
||||
|
||||
|
||||
def download_safetensors_index_file_from_hf(
|
||||
model_name_or_path: str,
|
||||
index_file: str,
|
||||
cache_dir: Optional[str],
|
||||
revision: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Download hf safetensors index file from Hugging Face Hub.
|
||||
|
||||
Args:
|
||||
model_name_or_path (str): The model name or path.
|
||||
cache_dir (Optional[str]): The cache directory to store the model
|
||||
weights. If None, will use HF defaults.
|
||||
revision (Optional[str]): The revision of the model.
|
||||
"""
|
||||
# Use file lock to prevent multiple processes from
|
||||
# downloading the same model weights at the same time.
|
||||
with get_lock(model_name_or_path, cache_dir):
|
||||
try:
|
||||
# Download the safetensors index file.
|
||||
hf_hub_download(
|
||||
repo_id=model_name_or_path,
|
||||
filename=index_file,
|
||||
cache_dir=cache_dir,
|
||||
revision=revision,
|
||||
local_files_only=huggingface_hub.constants.HF_HUB_OFFLINE,
|
||||
)
|
||||
# If file not found on remote or locally, we should not fail since
|
||||
# only some models will have index_file.
|
||||
except huggingface_hub.utils.EntryNotFoundError:
|
||||
logger.info("No %s found in remote.", index_file)
|
||||
except huggingface_hub.utils.LocalEntryNotFoundError:
|
||||
logger.info("No %s found in local cache.", index_file)
|
||||
|
||||
|
||||
# For models like Mistral-7B-v0.3, there are both sharded
|
||||
# safetensors files and a consolidated safetensors file.
|
||||
# Passing both of these to the weight loader functionality breaks.
|
||||
# So, we use the index_file to
|
||||
# look up which safetensors files should be used.
|
||||
def filter_duplicate_safetensors_files(hf_weights_files: List[str],
|
||||
hf_folder: str,
|
||||
index_file: str) -> List[str]:
|
||||
# model.safetensors.index.json is a mapping from keys in the
|
||||
# torch state_dict to safetensors file holding that weight.
|
||||
index_file_name = os.path.join(hf_folder, index_file)
|
||||
if not os.path.isfile(index_file_name):
|
||||
return hf_weights_files
|
||||
|
||||
# Iterate through the weight_map (weight_name: safetensors files)
|
||||
# to identify weights that we should use.
|
||||
with open(index_file_name) as f:
|
||||
weight_map = json.load(f)["weight_map"]
|
||||
weight_files_in_index = set()
|
||||
for weight_name in weight_map:
|
||||
weight_files_in_index.add(
|
||||
os.path.join(hf_folder, weight_map[weight_name]))
|
||||
# Filter out any fields that are not found in the index file.
|
||||
hf_weights_files = [
|
||||
f for f in hf_weights_files if f in weight_files_in_index
|
||||
]
|
||||
return hf_weights_files
|
||||
|
||||
|
||||
def filter_files_not_needed_for_inference(
|
||||
hf_weights_files: List[str]) -> List[str]:
|
||||
"""
|
||||
Exclude files that are not needed for inference.
|
||||
|
||||
See https://github.com/huggingface/transformers/blob/v4.34.0/src/transformers/trainer.py#L227-L233
|
||||
"""
|
||||
blacklist = [
|
||||
"training_args.bin",
|
||||
"optimizer.bin",
|
||||
"optimizer.pt",
|
||||
"scheduler.pt",
|
||||
"scaler.pt",
|
||||
]
|
||||
hf_weights_files = [
|
||||
f for f in hf_weights_files
|
||||
if not any(f.endswith(x) for x in blacklist)
|
||||
]
|
||||
return hf_weights_files
|
||||
|
||||
|
||||
# explicitly use pure text format, with a newline at the end
|
||||
# this makes it impossible to see the animation in the progress bar
|
||||
# but will avoid messing up with ray or multiprocessing, which wraps
|
||||
# each line of output with some prefix.
|
||||
_BAR_FORMAT = "{desc}: {percentage:3.0f}% Completed | {n_fmt}/{total_fmt} [{elapsed}<{remaining}, {rate_fmt}]\n" # noqa: E501
|
||||
|
||||
|
||||
def safetensors_weights_iterator(
|
||||
hf_weights_files: List[str]
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Iterate over the weights in the model safetensor files."""
|
||||
enable_tqdm = not torch.distributed.is_initialized(
|
||||
) or torch.distributed.get_rank() == 0
|
||||
for st_file in tqdm(
|
||||
hf_weights_files,
|
||||
desc="Loading safetensors checkpoint shards",
|
||||
disable=not enable_tqdm,
|
||||
bar_format=_BAR_FORMAT,
|
||||
):
|
||||
with safe_open(st_file, framework="pt") as f:
|
||||
for name in f.keys(): # noqa: SIM118
|
||||
param = f.get_tensor(name)
|
||||
yield name, param
|
||||
|
||||
|
||||
def pt_weights_iterator(
|
||||
hf_weights_files: List[str]
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Iterate over the weights in the model bin/pt files."""
|
||||
enable_tqdm = not torch.distributed.is_initialized(
|
||||
) or torch.distributed.get_rank() == 0
|
||||
for bin_file in tqdm(
|
||||
hf_weights_files,
|
||||
desc="Loading pt checkpoint shards",
|
||||
disable=not enable_tqdm,
|
||||
bar_format=_BAR_FORMAT,
|
||||
):
|
||||
state = torch.load(bin_file, map_location="cpu", weights_only=True)
|
||||
yield from state.items()
|
||||
del state
|
||||
|
||||
|
||||
def default_weight_loader(param: torch.Tensor,
|
||||
loaded_weight: torch.Tensor) -> None:
|
||||
"""Default weight loader."""
|
||||
try:
|
||||
if param.numel() == 1 and loaded_weight.numel() == 1:
|
||||
# Sometimes scalar values aren't considered tensors with shapes
|
||||
# so if both param and loaded_weight are a scalar,
|
||||
# "broadcast" instead of copy
|
||||
param.data.fill_(loaded_weight.item())
|
||||
else:
|
||||
assert param.size() == loaded_weight.size(), (
|
||||
f"Attempted to load weight ({loaded_weight.size()}) "
|
||||
f"into parameter ({param.size()})")
|
||||
|
||||
param.data.copy_(loaded_weight)
|
||||
except Exception:
|
||||
# NOTE: This exception is added for the purpose of setting breakpoint to
|
||||
# debug weight loading issues.
|
||||
raise
|
||||
|
||||
|
||||
def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> Optional[str]:
|
||||
"""Remap the name of FP8 k/v_scale parameters.
|
||||
|
||||
This function handles the remapping of FP8 k/v_scale parameter names.
|
||||
It detects if the given name ends with a suffix and attempts to remap
|
||||
it to the expected name format in the model. If the remapped name is not
|
||||
found in the params_dict, a warning is printed and None is returned.
|
||||
|
||||
Args:
|
||||
name (str): The original loaded checkpoint parameter name.
|
||||
params_dict (dict): Dictionary containing the model's named parameters.
|
||||
|
||||
Returns:
|
||||
str: The remapped parameter name if successful, or the original name
|
||||
if no remapping is needed.
|
||||
None: If the remapped name is not found in params_dict.
|
||||
"""
|
||||
if name.endswith(".kv_scale"):
|
||||
logger.warning_once(
|
||||
"DEPRECATED. Found kv_scale in the checkpoint. "
|
||||
"This format is deprecated in favor of separate k_scale and "
|
||||
"v_scale tensors and will be removed in a future release. "
|
||||
"Functionally, we will remap kv_scale to k_scale and duplicate "
|
||||
"k_scale to v_scale")
|
||||
# NOTE: we remap the deprecated kv_scale to k_scale
|
||||
remapped_name = name.replace(".kv_scale", ".attn.k_scale")
|
||||
if remapped_name not in params_dict:
|
||||
logger.warning_once(
|
||||
f"Found kv_scale in the checkpoint (e.g. {name}), "
|
||||
"but not found the expected name in the model "
|
||||
f"(e.g. {remapped_name}). kv_scale is "
|
||||
"not loaded.")
|
||||
return None
|
||||
return remapped_name
|
||||
|
||||
possible_scale_names = [".k_scale", ".v_scale"]
|
||||
modelopt_scale_names = [
|
||||
".self_attn.k_proj.k_scale", ".self_attn.v_proj.v_scale"
|
||||
]
|
||||
for scale_name in possible_scale_names:
|
||||
if name.endswith(scale_name):
|
||||
if any(mo_scale_name in name
|
||||
for mo_scale_name in modelopt_scale_names):
|
||||
remapped_name = name.replace(
|
||||
f".self_attn.{scale_name[1]}_proj{scale_name}",
|
||||
f".self_attn.attn{scale_name}")
|
||||
else:
|
||||
remapped_name = name.replace(scale_name, f".attn{scale_name}")
|
||||
if remapped_name not in params_dict:
|
||||
logger.warning_once(
|
||||
f"Found {scale_name} in the checkpoint (e.g. {name}), "
|
||||
"but not found the expected name in the model "
|
||||
f"(e.g. {remapped_name}). {scale_name} is "
|
||||
"not loaded.")
|
||||
return None
|
||||
return remapped_name
|
||||
|
||||
# If there were no matches, return the untouched param name
|
||||
return name
|
||||
@@ -0,0 +1,434 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/parameter.py
|
||||
|
||||
from fractions import Fraction
|
||||
from typing import Callable, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch.nn import Parameter
|
||||
|
||||
from fastvideo.v1.distributed import get_tensor_model_parallel_rank
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.utils import _make_synced_weight_loader
|
||||
|
||||
__all__ = [
|
||||
"BasevLLMParameter", "PackedvLLMParameter", "PerTensorScaleParameter",
|
||||
"ModelWeightParameter", "ChannelQuantScaleParameter",
|
||||
"GroupQuantScaleParameter", "PackedColumnParameter", "RowvLLMParameter"
|
||||
]
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class BasevLLMParameter(Parameter):
|
||||
"""
|
||||
Base parameter for vLLM linear layers. Extends the torch.nn.parameter
|
||||
by taking in a linear weight loader. Will copy the loaded weight
|
||||
into the parameter when the provided weight loader is called.
|
||||
"""
|
||||
|
||||
def __new__(cls, data: torch.Tensor, **kwargs):
|
||||
|
||||
return super().__new__(cls, data=data, requires_grad=False)
|
||||
|
||||
def __init__(self, data: torch.Tensor, weight_loader: Callable):
|
||||
"""
|
||||
Initialize the BasevLLMParameter
|
||||
|
||||
:param data: torch tensor with the parameter data
|
||||
:param weight_loader: weight loader callable
|
||||
|
||||
:returns: a torch.nn.parameter
|
||||
"""
|
||||
|
||||
# During weight loading, we often do something like:
|
||||
# narrowed_tensor = param.data.narrow(0, offset, len)
|
||||
# narrowed_tensor.copy_(real_weight)
|
||||
# expecting narrowed_tensor and param.data to share the same storage.
|
||||
# However, on TPUs, narrowed_tensor will lazily propagate to the base
|
||||
# tensor, which is param.data, leading to the redundant memory usage.
|
||||
# This sometimes causes OOM errors during model loading. To avoid this,
|
||||
# we sync the param tensor after its weight loader is called.
|
||||
from vllm.platforms import current_platform
|
||||
if current_platform.is_tpu():
|
||||
weight_loader = _make_synced_weight_loader(weight_loader)
|
||||
|
||||
self._weight_loader = weight_loader
|
||||
|
||||
@property
|
||||
def weight_loader(self):
|
||||
return self._weight_loader
|
||||
|
||||
def _is_1d_and_scalar(self, loaded_weight: torch.Tensor):
|
||||
cond1 = self.data.ndim == 1 and self.data.numel() == 1
|
||||
cond2 = loaded_weight.ndim == 0 and loaded_weight.numel() == 1
|
||||
return (cond1 and cond2)
|
||||
|
||||
def _assert_and_load(self, loaded_weight: torch.Tensor):
|
||||
assert (self.data.shape == loaded_weight.shape
|
||||
or self._is_1d_and_scalar(loaded_weight))
|
||||
self.data.copy_(loaded_weight)
|
||||
|
||||
def load_column_parallel_weight(self, loaded_weight: torch.Tensor):
|
||||
self._assert_and_load(loaded_weight)
|
||||
|
||||
def load_row_parallel_weight(self, loaded_weight: torch.Tensor):
|
||||
self._assert_and_load(loaded_weight)
|
||||
|
||||
def load_merged_column_weight(self, loaded_weight: torch.Tensor, **kwargs):
|
||||
self._assert_and_load(loaded_weight)
|
||||
|
||||
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs):
|
||||
self._assert_and_load(loaded_weight)
|
||||
|
||||
|
||||
class _ColumnvLLMParameter(BasevLLMParameter):
|
||||
"""
|
||||
Private class defining weight loading functionality
|
||||
(load_merged_column_weight, load_qkv_weight)
|
||||
for parameters being loaded into linear layers with column
|
||||
parallelism. This includes QKV and MLP layers which are
|
||||
not already fused on disk. Requires an output dimension
|
||||
to be defined. Called within the weight loader of
|
||||
each of the column parallel linear layers.
|
||||
"""
|
||||
|
||||
def __init__(self, output_dim: int, **kwargs):
|
||||
self._output_dim = output_dim
|
||||
super().__init__(**kwargs)
|
||||
|
||||
@property
|
||||
def output_dim(self):
|
||||
return self._output_dim
|
||||
|
||||
def load_column_parallel_weight(self, loaded_weight: torch.Tensor):
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
shard_size = self.data.shape[self.output_dim]
|
||||
loaded_weight = loaded_weight.narrow(self.output_dim,
|
||||
tp_rank * shard_size, shard_size)
|
||||
assert self.data.shape == loaded_weight.shape
|
||||
self.data.copy_(loaded_weight)
|
||||
|
||||
def load_merged_column_weight(self, loaded_weight: torch.Tensor, **kwargs):
|
||||
|
||||
shard_offset = kwargs.get("shard_offset")
|
||||
shard_size = kwargs.get("shard_size")
|
||||
if isinstance(
|
||||
self,
|
||||
(PackedColumnParameter,
|
||||
PackedvLLMParameter)) and self.packed_dim == self.output_dim:
|
||||
shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
|
||||
shard_offset=shard_offset, shard_size=shard_size)
|
||||
|
||||
param_data = self.data
|
||||
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
param_data = param_data.narrow(self.output_dim, shard_offset,
|
||||
shard_size)
|
||||
loaded_weight = loaded_weight.narrow(self.output_dim,
|
||||
tp_rank * shard_size, shard_size)
|
||||
assert param_data.shape == loaded_weight.shape
|
||||
param_data.copy_(loaded_weight)
|
||||
|
||||
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs):
|
||||
|
||||
shard_offset = kwargs.get("shard_offset")
|
||||
shard_size = kwargs.get("shard_size")
|
||||
shard_id = kwargs.get("shard_id")
|
||||
num_heads = kwargs.get("num_heads")
|
||||
|
||||
if isinstance(
|
||||
self,
|
||||
(PackedColumnParameter,
|
||||
PackedvLLMParameter)) and self.output_dim == self.packed_dim:
|
||||
shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
|
||||
shard_offset=shard_offset, shard_size=shard_size)
|
||||
|
||||
param_data = self.data
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
shard_id = tp_rank if shard_id == "q" else tp_rank // num_heads
|
||||
param_data = param_data.narrow(self.output_dim, shard_offset,
|
||||
shard_size)
|
||||
loaded_weight = loaded_weight.narrow(self.output_dim,
|
||||
shard_id * shard_size, shard_size)
|
||||
|
||||
assert param_data.shape == loaded_weight.shape
|
||||
param_data.copy_(loaded_weight)
|
||||
|
||||
|
||||
class RowvLLMParameter(BasevLLMParameter):
|
||||
"""
|
||||
Parameter class defining weight_loading functionality
|
||||
(load_row_parallel_weight) for parameters being loaded
|
||||
into linear layers with row parallel functionality.
|
||||
Requires an input_dim to be defined.
|
||||
"""
|
||||
|
||||
def __init__(self, input_dim: int, **kwargs):
|
||||
self._input_dim = input_dim
|
||||
super().__init__(**kwargs)
|
||||
|
||||
@property
|
||||
def input_dim(self):
|
||||
return self._input_dim
|
||||
|
||||
def load_row_parallel_weight(self, loaded_weight: torch.Tensor):
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
shard_size = self.data.shape[self.input_dim]
|
||||
loaded_weight = loaded_weight.narrow(self.input_dim,
|
||||
tp_rank * shard_size, shard_size)
|
||||
|
||||
if len(loaded_weight.shape) == 0:
|
||||
loaded_weight = loaded_weight.reshape(1)
|
||||
|
||||
assert self.data.shape == loaded_weight.shape
|
||||
self.data.copy_(loaded_weight)
|
||||
|
||||
|
||||
class ModelWeightParameter(_ColumnvLLMParameter, RowvLLMParameter):
|
||||
"""
|
||||
Parameter class for linear layer weights. Uses both column and
|
||||
row parallelism.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class GroupQuantScaleParameter(_ColumnvLLMParameter, RowvLLMParameter):
|
||||
"""
|
||||
Parameter class for weight scales loaded for weights with
|
||||
grouped quantization. Uses both column and row parallelism.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class ChannelQuantScaleParameter(_ColumnvLLMParameter):
|
||||
"""
|
||||
Parameter class for weight scales loaded for weights with
|
||||
channel-wise quantization. Equivalent to _ColumnvLLMParameter.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class PerTensorScaleParameter(BasevLLMParameter):
|
||||
"""
|
||||
Parameter class for scales where the number of scales is
|
||||
equivalent to the number of logical matrices in fused linear
|
||||
layers (e.g. for QKV, there are 3 scales loaded from disk).
|
||||
This is relevant to weights with per-tensor quantization.
|
||||
Adds functionality to map the scalers to a shard during
|
||||
weight loading.
|
||||
|
||||
Note: additional parameter manipulation may be handled
|
||||
for each quantization config specifically, within
|
||||
process_weights_after_loading
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.qkv_idxs = {"q": 0, "k": 1, "v": 2}
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def _shard_id_as_int(self, shard_id: Union[str, int]) -> int:
|
||||
if isinstance(shard_id, int):
|
||||
return shard_id
|
||||
|
||||
# if not int, assume shard_id for qkv
|
||||
# map to int and return
|
||||
assert isinstance(shard_id, str)
|
||||
assert shard_id in self.qkv_idxs
|
||||
return self.qkv_idxs[shard_id]
|
||||
|
||||
# For row parallel layers, no sharding needed
|
||||
# load weight into parameter as is
|
||||
def load_row_parallel_weight(self, *args, **kwargs):
|
||||
super().load_row_parallel_weight(*args, **kwargs)
|
||||
|
||||
def load_merged_column_weight(self, *args, **kwargs):
|
||||
self._load_into_shard_id(*args, **kwargs)
|
||||
|
||||
def load_qkv_weight(self, *args, **kwargs):
|
||||
self._load_into_shard_id(*args, **kwargs)
|
||||
|
||||
def load_column_parallel_weight(self, *args, **kwargs):
|
||||
super().load_row_parallel_weight(*args, **kwargs)
|
||||
|
||||
def _load_into_shard_id(self, loaded_weight: torch.Tensor,
|
||||
shard_id: Union[str, int], **kwargs):
|
||||
"""
|
||||
Slice the parameter data based on the shard id for
|
||||
loading.
|
||||
"""
|
||||
|
||||
param_data = self.data
|
||||
shard_id = self._shard_id_as_int(shard_id)
|
||||
|
||||
# AutoFP8 scales do not have a shape
|
||||
# compressed-tensors scales do have a shape
|
||||
if len(loaded_weight.shape) != 0:
|
||||
assert loaded_weight.shape[0] == 1
|
||||
loaded_weight = loaded_weight[0]
|
||||
|
||||
param_data = param_data[shard_id]
|
||||
assert param_data.shape == loaded_weight.shape
|
||||
param_data.copy_(loaded_weight)
|
||||
|
||||
|
||||
class PackedColumnParameter(_ColumnvLLMParameter):
|
||||
"""
|
||||
Parameter for model parameters which are packed on disk
|
||||
and support column parallelism only. See PackedvLLMParameter
|
||||
for more details on the packed properties.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
packed_factor: Union[int, Fraction],
|
||||
packed_dim: int,
|
||||
marlin_tile_size: Optional[int] = None,
|
||||
**kwargs):
|
||||
self._packed_factor = packed_factor
|
||||
self._packed_dim = packed_dim
|
||||
self._marlin_tile_size = marlin_tile_size
|
||||
super().__init__(**kwargs)
|
||||
|
||||
@property
|
||||
def packed_dim(self):
|
||||
return self._packed_dim
|
||||
|
||||
@property
|
||||
def packed_factor(self):
|
||||
return self._packed_factor
|
||||
|
||||
@property
|
||||
def marlin_tile_size(self):
|
||||
return self._marlin_tile_size
|
||||
|
||||
def adjust_shard_indexes_for_packing(self, shard_size, shard_offset):
|
||||
return _adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size,
|
||||
shard_offset=shard_offset,
|
||||
packed_factor=self.packed_factor,
|
||||
marlin_tile_size=self.marlin_tile_size)
|
||||
|
||||
|
||||
class PackedvLLMParameter(ModelWeightParameter):
|
||||
"""
|
||||
Parameter for model weights which are packed on disk.
|
||||
Example: GPTQ Marlin weights are int4 or int8, packed into int32.
|
||||
Extends the ModelWeightParameter to take in the
|
||||
packed factor, the packed dimension, and optionally, marlin
|
||||
tile size for marlin kernels. Adjusts the shard_size and
|
||||
shard_offset for fused linear layers model weight loading
|
||||
by accounting for packing and optionally, marlin tile size.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
packed_factor: Union[int, Fraction],
|
||||
packed_dim: int,
|
||||
marlin_tile_size: Optional[int] = None,
|
||||
**kwargs):
|
||||
self._packed_factor = packed_factor
|
||||
self._packed_dim = packed_dim
|
||||
self._marlin_tile_size = marlin_tile_size
|
||||
super().__init__(**kwargs)
|
||||
|
||||
@property
|
||||
def packed_dim(self):
|
||||
return self._packed_dim
|
||||
|
||||
@property
|
||||
def packed_factor(self):
|
||||
return self._packed_factor
|
||||
|
||||
@property
|
||||
def marlin_tile_size(self):
|
||||
return self._marlin_tile_size
|
||||
|
||||
def adjust_shard_indexes_for_packing(self, shard_size, shard_offset):
|
||||
return _adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size,
|
||||
shard_offset=shard_offset,
|
||||
packed_factor=self.packed_factor,
|
||||
marlin_tile_size=self.marlin_tile_size)
|
||||
|
||||
|
||||
class BlockQuantScaleParameter(_ColumnvLLMParameter, RowvLLMParameter):
|
||||
"""
|
||||
Parameter class for weight scales loaded for weights with
|
||||
block-wise quantization. Uses both column and row parallelism.
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
def permute_param_layout_(param: BasevLLMParameter, input_dim: int,
|
||||
output_dim: int, **kwargs) -> BasevLLMParameter:
|
||||
"""
|
||||
Permute a parameter's layout to the specified input and output dimensions,
|
||||
useful for forcing the parameter into a known layout, for example, if I need
|
||||
a packed (quantized) weight matrix to be in the layout
|
||||
{input_dim = 0, output_dim = 1, packed_dim = 0}
|
||||
then I can call:
|
||||
permute_param_layout_(x, input_dim=0, output_dim=1, packed_dim=0)
|
||||
to ensure x is in the correct layout (permuting it to the correct layout if
|
||||
required, asserting if it cannot get it to the correct layout)
|
||||
"""
|
||||
|
||||
curr_input_dim = getattr(param, "input_dim", None)
|
||||
curr_output_dim = getattr(param, "output_dim", None)
|
||||
|
||||
if curr_input_dim is None or curr_output_dim is None:
|
||||
assert param.data.dim() == 2,\
|
||||
"permute_param_layout_ only supports 2D parameters when either "\
|
||||
"input_dim or output_dim is not set"
|
||||
|
||||
# if one of the dimensions is not set, set it to the opposite of the other
|
||||
# we can only do this since we asserted the parameter is 2D above
|
||||
if curr_input_dim is None:
|
||||
assert curr_output_dim is not None,\
|
||||
"either input or output dim must be set"
|
||||
curr_input_dim = (curr_output_dim + 1) % 2
|
||||
if curr_output_dim is None:
|
||||
assert curr_input_dim is not None,\
|
||||
"either input or output dim must be set"
|
||||
curr_output_dim = (curr_input_dim + 1) % 2
|
||||
|
||||
# create permutation from the current layout to the layout with
|
||||
# self.input_dim at input_dim and self.output_dim at output_dim preserving
|
||||
# other dimensions
|
||||
perm = [
|
||||
i for i in range(param.data.dim())
|
||||
if i not in [curr_input_dim, curr_output_dim]
|
||||
]
|
||||
perm.insert(input_dim, curr_input_dim)
|
||||
perm.insert(output_dim, curr_output_dim)
|
||||
|
||||
if "packed_dim" in kwargs:
|
||||
assert hasattr(param, "packed_dim") and\
|
||||
param.packed_dim == perm[kwargs["packed_dim"]],\
|
||||
"permute_param_layout_ currently doesn't support repacking"
|
||||
|
||||
param.data = param.data.permute(*perm)
|
||||
if hasattr(param, "_input_dim"):
|
||||
param._input_dim = input_dim
|
||||
if hasattr(param, "_output_dim"):
|
||||
param._output_dim = output_dim
|
||||
if "packed_dim" in kwargs and hasattr(param, "_packed_dim"):
|
||||
param._packed_dim = kwargs["packed_dim"]
|
||||
|
||||
return param
|
||||
|
||||
|
||||
def _adjust_shard_indexes_for_marlin(shard_size, shard_offset,
|
||||
marlin_tile_size):
|
||||
return shard_size * marlin_tile_size, shard_offset * marlin_tile_size
|
||||
|
||||
|
||||
def _adjust_shard_indexes_for_packing(shard_size, shard_offset, packed_factor,
|
||||
marlin_tile_size):
|
||||
shard_size = shard_size // packed_factor
|
||||
shard_offset = shard_offset // packed_factor
|
||||
if marlin_tile_size is not None:
|
||||
return _adjust_shard_indexes_for_marlin(
|
||||
shard_size=shard_size,
|
||||
shard_offset=shard_offset,
|
||||
marlin_tile_size=marlin_tile_size)
|
||||
return shard_size, shard_offset
|
||||
@@ -0,0 +1,293 @@
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
import os
|
||||
import pickle
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from typing import AbstractSet, Callable, Dict, List, Optional, Tuple, Type, Union, TypeVar
|
||||
import importlib
|
||||
from functools import lru_cache
|
||||
import cloudpickle
|
||||
from torch import nn
|
||||
from fastvideo.v1.logger import logger
|
||||
|
||||
# huggingface class name: (component_name, fastvideo module name, fastvideo class name)
|
||||
_TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
|
||||
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
|
||||
}
|
||||
|
||||
_IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
|
||||
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
|
||||
}
|
||||
|
||||
_TEXT_ENCODER_MODELS = {
|
||||
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
|
||||
"LlamaModel": ("encoders", "llama", "LlamaModel"),
|
||||
}
|
||||
|
||||
_IMAGE_ENCODER_MODELS = {
|
||||
# "HunyuanVideoTransformer3DModel": ("image_encoder", "hunyuanvideo", "HunyuanVideoImageEncoder"),
|
||||
}
|
||||
|
||||
_VAE_MODELS = {
|
||||
"AutoencoderKLHunyuanVideo": ("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
|
||||
}
|
||||
|
||||
_FAST_VIDEO_MODELS = {
|
||||
**_TEXT_TO_VIDEO_DIT_MODELS,
|
||||
**_IMAGE_TO_VIDEO_DIT_MODELS,
|
||||
**_TEXT_ENCODER_MODELS,
|
||||
**_IMAGE_ENCODER_MODELS,
|
||||
**_VAE_MODELS,
|
||||
}
|
||||
|
||||
_SUBPROCESS_COMMAND = [
|
||||
sys.executable, "-m", "fastvideo.v1.models.dits.registry"
|
||||
]
|
||||
|
||||
|
||||
_T = TypeVar("_T")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ModelInfo:
|
||||
architecture: str
|
||||
|
||||
@staticmethod
|
||||
def from_model_cls(model: Type[nn.Module]) -> "_ModelInfo":
|
||||
return _ModelInfo(
|
||||
architecture=model.__name__,)
|
||||
|
||||
|
||||
class _BaseRegisteredModel(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def inspect_model_cls(self) -> _ModelInfo:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def load_model_cls(self) -> Type[nn.Module]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _RegisteredModel(_BaseRegisteredModel):
|
||||
"""
|
||||
Represents a model that has already been imported in the main process.
|
||||
"""
|
||||
|
||||
interfaces: _ModelInfo
|
||||
model_cls: Type[nn.Module]
|
||||
|
||||
@staticmethod
|
||||
def from_model_cls(model_cls: Type[nn.Module]):
|
||||
return _RegisteredModel(
|
||||
interfaces=_ModelInfo.from_model_cls(model_cls),
|
||||
model_cls=model_cls,
|
||||
)
|
||||
|
||||
def inspect_model_cls(self) -> _ModelInfo:
|
||||
return self.interfaces
|
||||
|
||||
def load_model_cls(self) -> Type[nn.Module]:
|
||||
return self.model_cls
|
||||
|
||||
def _run_in_subprocess(fn: Callable[[], _T]) -> _T:
|
||||
# NOTE: We use a temporary directory instead of a temporary file to avoid
|
||||
# issues like https://stackoverflow.com/questions/23212435/permission-denied-to-write-to-my-temporary-file
|
||||
with tempfile.TemporaryDirectory() as tempdir:
|
||||
output_filepath = os.path.join(tempdir, "registry_output.tmp")
|
||||
|
||||
# `cloudpickle` allows pickling lambda functions directly
|
||||
input_bytes = cloudpickle.dumps((fn, output_filepath))
|
||||
|
||||
# cannot use `sys.executable __file__` here because the script
|
||||
# contains relative imports
|
||||
returned = subprocess.run(_SUBPROCESS_COMMAND,
|
||||
input=input_bytes,
|
||||
capture_output=True)
|
||||
|
||||
# check if the subprocess is successful
|
||||
try:
|
||||
returned.check_returncode()
|
||||
except Exception as e:
|
||||
# wrap raised exception to provide more information
|
||||
raise RuntimeError(f"Error raised in subprocess:\n"
|
||||
f"{returned.stderr.decode()}") from e
|
||||
|
||||
with open(output_filepath, "rb") as f:
|
||||
return pickle.load(f)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _LazyRegisteredModel(_BaseRegisteredModel):
|
||||
"""
|
||||
Represents a model that has not been imported in the main process.
|
||||
"""
|
||||
module_name: str
|
||||
component_name: str
|
||||
class_name: str
|
||||
|
||||
# Performed in another process to avoid initializing CUDA
|
||||
def inspect_model_cls(self) -> _ModelInfo:
|
||||
return _run_in_subprocess(
|
||||
lambda: _ModelInfo.from_model_cls(self.load_model_cls()))
|
||||
|
||||
def load_model_cls(self) -> Type[nn.Module]:
|
||||
mod = importlib.import_module(self.module_name)
|
||||
return getattr(mod, self.class_name)
|
||||
|
||||
|
||||
@lru_cache(maxsize=128)
|
||||
def _try_load_model_cls(
|
||||
model_arch: str,
|
||||
model: _BaseRegisteredModel,
|
||||
) -> Optional[Type[nn.Module]]:
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
current_platform.verify_model_arch(model_arch)
|
||||
try:
|
||||
return model.load_model_cls()
|
||||
except Exception:
|
||||
logger.exception("Error in loading model architecture '%s'",
|
||||
model_arch)
|
||||
return None
|
||||
|
||||
|
||||
@lru_cache(maxsize=128)
|
||||
def _try_inspect_model_cls(
|
||||
model_arch: str,
|
||||
model: _BaseRegisteredModel,
|
||||
) -> Optional[_ModelInfo]:
|
||||
try:
|
||||
return model.inspect_model_cls()
|
||||
except Exception:
|
||||
logger.exception("Error in inspecting model architecture '%s'",
|
||||
model_arch)
|
||||
return None
|
||||
|
||||
|
||||
@dataclass
|
||||
class _ModelRegistry:
|
||||
# Keyed by model_arch
|
||||
models: Dict[str, _BaseRegisteredModel] = field(default_factory=dict)
|
||||
|
||||
def get_supported_archs(self) -> AbstractSet[str]:
|
||||
return self.models.keys()
|
||||
|
||||
def register_model(
|
||||
self,
|
||||
model_arch: str,
|
||||
model_cls: Union[Type[nn.Module], str],
|
||||
) -> None:
|
||||
"""
|
||||
Register an external model to be used in vLLM.
|
||||
|
||||
:code:`model_cls` can be either:
|
||||
|
||||
- A :class:`torch.nn.Module` class directly referencing the model.
|
||||
- A string in the format :code:`<module>:<class>` which can be used to
|
||||
lazily import the model. This is useful to avoid initializing CUDA
|
||||
when importing the model and thus the related error
|
||||
:code:`RuntimeError: Cannot re-initialize CUDA in forked subprocess`.
|
||||
"""
|
||||
if model_arch in self.models:
|
||||
logger.warning(
|
||||
"Model architecture %s is already registered, and will be "
|
||||
"overwritten by the new model class %s.", model_arch,
|
||||
model_cls)
|
||||
|
||||
if isinstance(model_cls, str):
|
||||
split_str = model_cls.split(":")
|
||||
if len(split_str) != 2:
|
||||
msg = "Expected a string in the format `<module>:<class>`"
|
||||
raise ValueError(msg)
|
||||
|
||||
model = _LazyRegisteredModel(*split_str)
|
||||
else:
|
||||
model = _RegisteredModel.from_model_cls(model_cls)
|
||||
|
||||
self.models[model_arch] = model
|
||||
|
||||
def _raise_for_unsupported(self, architectures: List[str]):
|
||||
all_supported_archs = self.get_supported_archs()
|
||||
|
||||
if any(arch in all_supported_archs for arch in architectures):
|
||||
raise ValueError(
|
||||
f"Model architectures {architectures} failed "
|
||||
"to be inspected. Please check the logs for more details.")
|
||||
|
||||
raise ValueError(
|
||||
f"Model architectures {architectures} are not supported for now. "
|
||||
f"Supported architectures: {all_supported_archs}")
|
||||
|
||||
def _try_load_model_cls(self,
|
||||
model_arch: str) -> Optional[Type[nn.Module]]:
|
||||
if model_arch not in self.models:
|
||||
return None
|
||||
|
||||
return _try_load_model_cls(model_arch, self.models[model_arch])
|
||||
|
||||
def _try_inspect_model_cls(self, model_arch: str) -> Optional[_ModelInfo]:
|
||||
if model_arch not in self.models:
|
||||
return None
|
||||
|
||||
return _try_inspect_model_cls(model_arch, self.models[model_arch])
|
||||
|
||||
def _normalize_archs(
|
||||
self,
|
||||
architectures: Union[str, List[str]],
|
||||
) -> List[str]:
|
||||
if isinstance(architectures, str):
|
||||
architectures = [architectures]
|
||||
if not architectures:
|
||||
logger.warning("No model architectures are specified")
|
||||
|
||||
normalized_arch = []
|
||||
for model in architectures:
|
||||
if model not in self.models:
|
||||
model = "TransformersModel"
|
||||
normalized_arch.append(model)
|
||||
return normalized_arch
|
||||
|
||||
def inspect_model_cls(
|
||||
self,
|
||||
architectures: Union[str, List[str]],
|
||||
) -> Tuple[_ModelInfo, str]:
|
||||
architectures = self._normalize_archs(architectures)
|
||||
|
||||
for arch in architectures:
|
||||
model_info = self._try_inspect_model_cls(arch)
|
||||
if model_info is not None:
|
||||
return (model_info, arch)
|
||||
|
||||
return self._raise_for_unsupported(architectures)
|
||||
|
||||
def resolve_model_cls(
|
||||
self,
|
||||
architectures: Union[str, List[str]],
|
||||
) -> Tuple[Type[nn.Module], str]:
|
||||
architectures = self._normalize_archs(architectures)
|
||||
|
||||
for arch in architectures:
|
||||
model_cls = self._try_load_model_cls(arch)
|
||||
if model_cls is not None:
|
||||
return (model_cls, arch)
|
||||
|
||||
return self._raise_for_unsupported(architectures)
|
||||
|
||||
|
||||
|
||||
|
||||
ModelRegistry = _ModelRegistry({
|
||||
model_arch:
|
||||
_LazyRegisteredModel(
|
||||
module_name=f"fastvideo.v1.models.{component_name}.{mod_relname}",
|
||||
component_name=component_name,
|
||||
class_name=cls_name,
|
||||
)
|
||||
for model_arch, (component_name, mod_relname, cls_name) in _FAST_VIDEO_MODELS.items()
|
||||
})
|
||||
@@ -0,0 +1,239 @@
|
||||
# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
#
|
||||
# Modified from diffusers==0.29.2
|
||||
#
|
||||
# ==============================================================================
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
|
||||
"""
|
||||
Output class for the scheduler's `step` function output.
|
||||
|
||||
Args:
|
||||
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
|
||||
denoising loop.
|
||||
"""
|
||||
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
|
||||
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
"""
|
||||
Euler scheduler.
|
||||
|
||||
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
|
||||
methods the library implements for all schedulers such as loading and saving.
|
||||
|
||||
Args:
|
||||
num_train_timesteps (`int`, defaults to 1000):
|
||||
The number of diffusion steps to train the model.
|
||||
timestep_spacing (`str`, defaults to `"linspace"`):
|
||||
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
|
||||
shift (`float`, defaults to 1.0):
|
||||
The shift value for the timestep schedule.
|
||||
reverse (`bool`, defaults to `True`):
|
||||
Whether to reverse the timestep schedule.
|
||||
"""
|
||||
|
||||
_compatibles = []
|
||||
order = 1
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_train_timesteps: int = 1000,
|
||||
shift: float = 1.0,
|
||||
reverse: bool = True,
|
||||
solver: str = "euler",
|
||||
n_tokens: Optional[int] = None,
|
||||
):
|
||||
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
|
||||
|
||||
if not reverse:
|
||||
sigmas = sigmas.flip(0)
|
||||
|
||||
self.sigmas = sigmas
|
||||
# the value fed to model
|
||||
self.timesteps = (sigmas[:-1] * num_train_timesteps).to(dtype=torch.float32)
|
||||
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
self.supported_solver = ["euler"]
|
||||
if solver not in self.supported_solver:
|
||||
raise ValueError(f"Solver {solver} not supported. Supported solvers: {self.supported_solver}")
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
"""
|
||||
The index counter for current timestep. It will increase 1 after each scheduler step.
|
||||
"""
|
||||
return self._step_index
|
||||
|
||||
@property
|
||||
def begin_index(self):
|
||||
"""
|
||||
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
|
||||
"""
|
||||
return self._begin_index
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
|
||||
def set_begin_index(self, begin_index: int = 0):
|
||||
"""
|
||||
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
|
||||
|
||||
Args:
|
||||
begin_index (`int`):
|
||||
The begin index for the scheduler.
|
||||
"""
|
||||
self._begin_index = begin_index
|
||||
|
||||
def _sigma_to_t(self, sigma):
|
||||
return sigma * self.config.num_train_timesteps
|
||||
|
||||
def set_timesteps(
|
||||
self,
|
||||
num_inference_steps: int,
|
||||
device: Union[str, torch.device] = None,
|
||||
n_tokens: int = None,
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
Args:
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
n_tokens (`int`, *optional*):
|
||||
Number of tokens in the input sequence.
|
||||
"""
|
||||
self.num_inference_steps = num_inference_steps
|
||||
|
||||
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
|
||||
sigmas = self.sd3_time_shift(sigmas)
|
||||
|
||||
if not self.config.reverse:
|
||||
sigmas = 1 - sigmas
|
||||
|
||||
self.sigmas = sigmas
|
||||
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(dtype=torch.float32, device=device)
|
||||
|
||||
# Reset step index
|
||||
self._step_index = None
|
||||
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None):
|
||||
if schedule_timesteps is None:
|
||||
schedule_timesteps = self.timesteps
|
||||
|
||||
indices = (schedule_timesteps == timestep).nonzero()
|
||||
|
||||
# The sigma index that is taken for the **very** first `step`
|
||||
# is always the second index (or the last index if there is only 1)
|
||||
# This way we can ensure we don't accidentally skip a sigma in
|
||||
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||
pos = 1 if len(indices) > 1 else 0
|
||||
|
||||
return indices[pos].item()
|
||||
|
||||
def _init_step_index(self, timestep):
|
||||
if self.begin_index is None:
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
timestep = timestep.to(self.timesteps.device)
|
||||
self._step_index = self.index_for_timestep(timestep)
|
||||
else:
|
||||
self._step_index = self._begin_index
|
||||
|
||||
def scale_model_input(self, sample: torch.Tensor, timestep: Optional[int] = None) -> torch.Tensor:
|
||||
return sample
|
||||
|
||||
def sd3_time_shift(self, t: torch.Tensor):
|
||||
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
return_dict: bool = True,
|
||||
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||
process from the learned model outputs (most often the predicted noise).
|
||||
|
||||
Args:
|
||||
model_output (`torch.FloatTensor`):
|
||||
The direct output from learned diffusion model.
|
||||
timestep (`float`):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.FloatTensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
generator (`torch.Generator`, *optional*):
|
||||
A random number generator.
|
||||
n_tokens (`int`, *optional*):
|
||||
Number of tokens in the input sequence.
|
||||
return_dict (`bool`):
|
||||
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
|
||||
tuple.
|
||||
|
||||
Returns:
|
||||
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
|
||||
returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||
"""
|
||||
|
||||
if (isinstance(timestep, int) or isinstance(timestep, torch.IntTensor)
|
||||
or isinstance(timestep, torch.LongTensor)):
|
||||
raise ValueError(("Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
" one of the `scheduler.timesteps` as a timestep."), )
|
||||
|
||||
if self.step_index is None:
|
||||
self._init_step_index(timestep)
|
||||
|
||||
# Upcast to avoid precision issues when computing prev_sample
|
||||
sample = sample.to(torch.float32)
|
||||
|
||||
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
|
||||
|
||||
if self.config.solver == "euler":
|
||||
prev_sample = sample + model_output.to(torch.float32) * dt
|
||||
else:
|
||||
raise ValueError(f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}")
|
||||
|
||||
# upon completion increase step index by one
|
||||
self._step_index += 1
|
||||
|
||||
if not return_dict:
|
||||
return (prev_sample, )
|
||||
|
||||
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
@@ -0,0 +1,197 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from transformers.utils import ModelOutput
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
def use_default(value, default):
|
||||
return value if value is not None else default
|
||||
|
||||
@dataclass
|
||||
class TextEncoderModelOutput(ModelOutput):
|
||||
"""
|
||||
Base class for model's outputs that also contains a pooling of the last hidden states.
|
||||
|
||||
Args:
|
||||
hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
|
||||
Sequence of hidden-states at the output of the last layer of the model.
|
||||
attention_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
|
||||
Mask to avoid performing attention on padding token indices. Mask values selected in ``[0, 1]``:
|
||||
hidden_states_list (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed):
|
||||
Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
|
||||
one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
|
||||
Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
|
||||
text_outputs (`list`, *optional*, returned when `return_texts=True` is passed):
|
||||
List of decoded texts.
|
||||
"""
|
||||
|
||||
hidden_state: torch.FloatTensor = None
|
||||
attention_mask: Optional[torch.LongTensor] = None
|
||||
text_outputs: Optional[list] = None
|
||||
|
||||
|
||||
class TextEncoder(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
text_encoder,
|
||||
tokenizer,
|
||||
max_length: int,
|
||||
text_encoder_precision: Optional[str] = None,
|
||||
text_encoder_path: Optional[str] = None,
|
||||
output_key: Optional[str] = None,
|
||||
use_attention_mask: bool = True,
|
||||
prompt_template: Optional[dict] = None,
|
||||
prompt_template_video: Optional[dict] = None,
|
||||
hidden_state_skip_layer: Optional[int] = None,
|
||||
apply_final_norm: bool = False,
|
||||
device=None,
|
||||
):
|
||||
super().__init__()
|
||||
# TODO(will): check if there's a cleaner way to do this
|
||||
self.text_encoder_type = text_encoder.config.architectures[0]
|
||||
self.max_length = max_length
|
||||
self.precision = text_encoder_precision
|
||||
self.model_path = text_encoder_path
|
||||
self.use_attention_mask = use_attention_mask
|
||||
if prompt_template_video is not None:
|
||||
assert (use_attention_mask is True), "Attention mask is True required when training videos."
|
||||
self.prompt_template = prompt_template
|
||||
self.prompt_template_video = prompt_template_video
|
||||
self.hidden_state_skip_layer = hidden_state_skip_layer
|
||||
self.apply_final_norm = apply_final_norm
|
||||
|
||||
|
||||
if "T5" in self.text_encoder_type:
|
||||
self.output_key = output_key or "last_hidden_state"
|
||||
elif "CLIPTextModel" in self.text_encoder_type:
|
||||
self.output_key = output_key or "pooler_output"
|
||||
elif "LlamaModel" in self.text_encoder_type or "glm" in self.text_encoder_type:
|
||||
self.output_key = output_key or "last_hidden_state"
|
||||
else:
|
||||
raise ValueError(f"Unsupported text encoder type: {self.text_encoder_type}")
|
||||
|
||||
self.model = text_encoder
|
||||
# self.dtype = self.model.dtype
|
||||
self.device = device
|
||||
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def __repr__(self):
|
||||
return f"{self.text_encoder_type} ({self.precision} - {self.model_path})"
|
||||
|
||||
@staticmethod
|
||||
def apply_text_to_template(text, template, prevent_empty_text=True):
|
||||
"""
|
||||
Apply text to template.
|
||||
|
||||
Args:
|
||||
text (str): Input text.
|
||||
template (str or list): Template string or list of chat conversation.
|
||||
prevent_empty_text (bool): If True, we will prevent the user text from being empty
|
||||
by adding a space. Defaults to True.
|
||||
"""
|
||||
if isinstance(template, str):
|
||||
# Will send string to tokenizer. Used for llm
|
||||
return template.format(text)
|
||||
else:
|
||||
raise TypeError(f"Unsupported template type: {type(template)}")
|
||||
|
||||
def text2tokens(self, text):
|
||||
"""
|
||||
Tokenize the input text.
|
||||
|
||||
Args:
|
||||
text (str or list): Input text.
|
||||
"""
|
||||
if self.prompt_template_video is not None:
|
||||
prompt_template = self.prompt_template_video["template"]
|
||||
|
||||
text = self.apply_text_to_template(text, prompt_template)
|
||||
|
||||
|
||||
kwargs = dict(
|
||||
truncation=True,
|
||||
max_length=self.max_length,
|
||||
padding="max_length",
|
||||
return_tensors="pt",
|
||||
)
|
||||
return self.tokenizer(
|
||||
text,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_attention_mask=True,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def encode(
|
||||
self,
|
||||
batch_encoding,
|
||||
use_attention_mask=None,
|
||||
hidden_state_skip_layer=None,
|
||||
device=None,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
batch_encoding (dict): Batch encoding from tokenizer.
|
||||
use_attention_mask (bool): Whether to use attention mask. If None, use self.use_attention_mask.
|
||||
Defaults to None.
|
||||
output_hidden_states (bool): Whether to output hidden states. If False, return the value of
|
||||
self.output_key. If True, return the entire output. If set self.hidden_state_skip_layer,
|
||||
output_hidden_states will be set True. Defaults to False.
|
||||
hidden_state_skip_layer (int): Number of hidden states to hidden_state_skip_layer. 0 means the last layer.
|
||||
If None, self.output_key will be used. Defaults to None.
|
||||
return_texts (bool): Whether to return the decoded texts. Defaults to False.
|
||||
"""
|
||||
device = self.model.device if device is None else device
|
||||
use_attention_mask = use_default(use_attention_mask, self.use_attention_mask)
|
||||
hidden_state_skip_layer = use_default(hidden_state_skip_layer, self.hidden_state_skip_layer)
|
||||
attention_mask = (batch_encoding["attention_mask"].to(device) if use_attention_mask else None)
|
||||
|
||||
# note: clip will need attention mask
|
||||
outputs = self.model(
|
||||
input_ids=batch_encoding["input_ids"].to(device),
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=hidden_state_skip_layer is not None,
|
||||
)
|
||||
if hidden_state_skip_layer is not None:
|
||||
last_hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
|
||||
# Real last hidden state already has layer norm applied. So here we only apply it
|
||||
# for intermediate layers.
|
||||
if hidden_state_skip_layer > 0 and self.apply_final_norm:
|
||||
last_hidden_state = self.model.final_layer_norm(last_hidden_state)
|
||||
else:
|
||||
last_hidden_state = outputs[self.output_key]
|
||||
|
||||
# Remove hidden states of instruction tokens, only keep prompt tokens.
|
||||
if self.prompt_template_video is not None:
|
||||
|
||||
crop_start = self.prompt_template_video.get("crop_start", -1)
|
||||
|
||||
last_hidden_state = last_hidden_state[:, crop_start:]
|
||||
attention_mask = (attention_mask[:, crop_start:] if use_attention_mask else None)
|
||||
total_length = attention_mask.sum()
|
||||
last_hidden_state = last_hidden_state[:, :total_length]
|
||||
return TextEncoderModelOutput(last_hidden_state)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
text,
|
||||
use_attention_mask=None,
|
||||
output_hidden_states=False,
|
||||
hidden_state_skip_layer=None,
|
||||
return_texts=False,
|
||||
):
|
||||
batch_encoding = self.text2tokens(text)
|
||||
return self.encode(
|
||||
batch_encoding,
|
||||
use_attention_mask=use_attention_mask,
|
||||
output_hidden_states=output_hidden_states,
|
||||
hidden_state_skip_layer=hidden_state_skip_layer,
|
||||
return_texts=return_texts,
|
||||
)
|
||||
@@ -0,0 +1,97 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/utils.py
|
||||
"""Utils for model executor."""
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import torch
|
||||
|
||||
# TODO(PY): move it elsewhere
|
||||
def auto_attributes(init_func):
|
||||
"""
|
||||
Decorator that automatically adds all initialization arguments as object attributes.
|
||||
|
||||
Example:
|
||||
@auto_attributes
|
||||
def __init__(self, a=1, b=2):
|
||||
pass
|
||||
|
||||
# This will automatically set:
|
||||
# - self.a = 1 and self.b = 2
|
||||
# - self.config.a = 1 and self.config.b = 2
|
||||
"""
|
||||
def wrapper(self, *args, **kwargs):
|
||||
# Get the function signature
|
||||
import inspect
|
||||
signature = inspect.signature(init_func)
|
||||
parameters = signature.parameters
|
||||
|
||||
# Get parameter names (excluding 'self')
|
||||
param_names = list(parameters.keys())[1:]
|
||||
|
||||
# Bind arguments to parameters
|
||||
bound_args = signature.bind(self, *args, **kwargs)
|
||||
bound_args.apply_defaults()
|
||||
|
||||
# Create config object if it doesn't exist
|
||||
if not hasattr(self, 'config'):
|
||||
self.config = type('Config', (), {})()
|
||||
|
||||
# Set attributes on self and self.config
|
||||
for name in param_names:
|
||||
if name in bound_args.arguments:
|
||||
value = bound_args.arguments[name]
|
||||
setattr(self, name, value)
|
||||
setattr(self.config, name, value)
|
||||
|
||||
# Call the original __init__ function
|
||||
return init_func(self, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
def set_random_seed(seed: int) -> None:
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
current_platform.seed_everything(seed)
|
||||
|
||||
|
||||
def set_weight_attrs(
|
||||
weight: torch.Tensor,
|
||||
weight_attrs: Optional[Dict[str, Any]],
|
||||
):
|
||||
"""Set attributes on a weight tensor.
|
||||
|
||||
This method is used to set attributes on a weight tensor. This method
|
||||
will not overwrite existing attributes.
|
||||
|
||||
Args:
|
||||
weight: The weight tensor.
|
||||
weight_attrs: A dictionary of attributes to set on the weight tensor.
|
||||
"""
|
||||
if weight_attrs is None:
|
||||
return
|
||||
for key, value in weight_attrs.items():
|
||||
assert not hasattr(
|
||||
weight, key), (f"Overwriting existing tensor attribute: {key}")
|
||||
|
||||
# NOTE(woosuk): During weight loading, we often do something like:
|
||||
# narrowed_tensor = param.data.narrow(0, offset, len)
|
||||
# narrowed_tensor.copy_(real_weight)
|
||||
# expecting narrowed_tensor and param.data to share the same storage.
|
||||
# However, on TPUs, narrowed_tensor will lazily propagate to the base
|
||||
# tensor, which is param.data, leading to the redundant memory usage.
|
||||
# This sometimes causes OOM errors during model loading. To avoid this,
|
||||
# we sync the param tensor after its weight loader is called.
|
||||
# TODO(woosuk): Remove this hack once we have a better solution.
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
if current_platform.is_tpu() and key == "weight_loader":
|
||||
value = _make_synced_weight_loader(value)
|
||||
setattr(weight, key, value)
|
||||
|
||||
|
||||
def _make_synced_weight_loader(original_weight_loader):
|
||||
|
||||
def _synced_weight_loader(param, *args, **kwargs):
|
||||
original_weight_loader(param, *args, **kwargs)
|
||||
torch._sync(param)
|
||||
|
||||
return _synced_weight_loader
|
||||
@@ -0,0 +1,469 @@
|
||||
import torch
|
||||
from typing import Optional, Tuple
|
||||
import numpy as np
|
||||
import torch.nn as nn
|
||||
from abc import abstractmethod, ABC
|
||||
import torch.distributed as dist
|
||||
from fastvideo.v1.distributed import tensor_model_parallel_all_gather, get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size
|
||||
from math import prod
|
||||
class ParallelTiledVAE(ABC):
|
||||
def __init__(self, *args, **kwargs):
|
||||
# Check if subclass has defined all required properties
|
||||
required_attributes = [
|
||||
'tile_sample_min_height',
|
||||
'tile_sample_min_width',
|
||||
'tile_sample_min_num_frames',
|
||||
'tile_sample_stride_height',
|
||||
'tile_sample_stride_width',
|
||||
'tile_sample_stride_num_frames',
|
||||
'use_tiling'
|
||||
]
|
||||
|
||||
for attr in required_attributes:
|
||||
if not hasattr(self, attr):
|
||||
raise AttributeError(f"Subclasses of ParallelVAE must define '{attr}' property")
|
||||
|
||||
@abstractmethod
|
||||
def _encode(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def _decode(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
|
||||
def encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, num_channels, num_frames, height, width = x.shape
|
||||
|
||||
if self.use_tiling and num_frames > self.tile_sample_min_num_frames:
|
||||
latents = self.tiled_encode(x)
|
||||
elif self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height):
|
||||
latents = self.spatial_tiled_encode(x)
|
||||
else:
|
||||
latents = self._encode(x)
|
||||
return DiagonalGaussianDistribution(latents)
|
||||
|
||||
|
||||
|
||||
def decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, num_channels, num_frames, height, width = z.shape
|
||||
tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio
|
||||
tile_latent_min_width = self.tile_sample_stride_width // self.spatial_compression_ratio
|
||||
tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio
|
||||
|
||||
if self.use_tiling and get_sequence_model_parallel_world_size() > 1:
|
||||
return self.parallel_tiled_decode(z)
|
||||
if self.use_tiling and num_frames > tile_latent_min_num_frames:
|
||||
return self.tiled_decode(z)
|
||||
|
||||
if self.use_tiling and (width > tile_latent_min_width or height > tile_latent_min_height):
|
||||
return self.spatial_tiled_decode(z)
|
||||
|
||||
return self._decode(z)
|
||||
|
||||
|
||||
def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[-2], b.shape[-2], blend_extent)
|
||||
for y in range(blend_extent):
|
||||
b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * (
|
||||
y / blend_extent
|
||||
)
|
||||
return b
|
||||
|
||||
def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[-1], b.shape[-1], blend_extent)
|
||||
for x in range(blend_extent):
|
||||
b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * (
|
||||
x / blend_extent
|
||||
)
|
||||
return b
|
||||
|
||||
def blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[-3], b.shape[-3], blend_extent)
|
||||
for x in range(blend_extent):
|
||||
b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (1 - x / blend_extent) + b[:, :, x, :, :] * (
|
||||
x / blend_extent
|
||||
)
|
||||
return b
|
||||
|
||||
def spatial_tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
r"""Encode a batch of images using a tiled encoder.
|
||||
|
||||
Args:
|
||||
x (`torch.Tensor`): Input batch of videos.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
The latent representation of the encoded videos.
|
||||
"""
|
||||
batch_size, num_channels, num_frames, height, width = x.shape
|
||||
latent_height = height // self.spatial_compression_ratio
|
||||
latent_width = width // self.spatial_compression_ratio
|
||||
|
||||
tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio
|
||||
tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio
|
||||
tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio
|
||||
tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio
|
||||
|
||||
blend_height = tile_latent_min_height - tile_latent_stride_height
|
||||
blend_width = tile_latent_min_width - tile_latent_stride_width
|
||||
|
||||
# Split x into overlapping tiles and encode them separately.
|
||||
# The tiles have an overlap to avoid seams between tiles.
|
||||
rows = []
|
||||
for i in range(0, height, self.tile_sample_stride_height):
|
||||
row = []
|
||||
for j in range(0, width, self.tile_sample_stride_width):
|
||||
tile = x[:, :, :, i : i + self.tile_sample_min_height, j : j + self.tile_sample_min_width]
|
||||
tile = self._encode(tile)
|
||||
row.append(tile)
|
||||
rows.append(row)
|
||||
|
||||
return self._merge_spatial_tiles(rows, blend_height, blend_width,tile_latent_stride_height, tile_latent_stride_width)
|
||||
|
||||
def _parallel_data_generator(self, gathered_results,
|
||||
gathered_dim_metadata):
|
||||
global_idx = 0
|
||||
for i, per_rank_metadata in enumerate(gathered_dim_metadata):
|
||||
_start_shape = 0
|
||||
for shape in per_rank_metadata:
|
||||
mul_shape = prod(shape)
|
||||
yield (gathered_results[i, _start_shape:_start_shape +
|
||||
mul_shape].reshape(shape), global_idx)
|
||||
_start_shape += mul_shape
|
||||
global_idx += 1
|
||||
|
||||
def parallel_tiled_decode(self, z: torch.FloatTensor) -> torch.FloatTensor:
|
||||
"""
|
||||
Parallel version of tiled_decode that distributes both temporal and spatial computation across GPUs
|
||||
"""
|
||||
world_size, rank = get_sequence_model_parallel_world_size(), get_sequence_model_parallel_rank()
|
||||
B, C, T, H, W = z.shape
|
||||
|
||||
# Calculate parameters
|
||||
tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio
|
||||
tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio
|
||||
tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio
|
||||
tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio
|
||||
tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio
|
||||
tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio
|
||||
|
||||
blend_height = self.tile_sample_min_height - self.tile_sample_stride_height
|
||||
blend_width = self.tile_sample_min_width - self.tile_sample_stride_width
|
||||
blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
|
||||
|
||||
# Calculate tile dimensions
|
||||
num_t_tiles = (T + tile_latent_stride_num_frames - 1) // tile_latent_stride_num_frames
|
||||
num_h_tiles = (H + tile_latent_stride_height - 1) // tile_latent_stride_height
|
||||
num_w_tiles = (W + tile_latent_stride_width - 1) // tile_latent_stride_width
|
||||
total_spatial_tiles = num_h_tiles * num_w_tiles
|
||||
total_tiles = num_t_tiles * total_spatial_tiles
|
||||
|
||||
# Calculate tiles per rank and padding
|
||||
tiles_per_rank = (total_tiles + world_size - 1) // world_size
|
||||
start_tile_idx = rank * tiles_per_rank
|
||||
end_tile_idx = min((rank + 1) * tiles_per_rank, total_tiles)
|
||||
|
||||
local_results = []
|
||||
local_dim_metadata = []
|
||||
# Process assigned tiles
|
||||
for local_idx, global_idx in enumerate(range(start_tile_idx, end_tile_idx)):
|
||||
t_idx = global_idx // total_spatial_tiles
|
||||
spatial_idx = global_idx % total_spatial_tiles
|
||||
h_idx = spatial_idx // num_w_tiles
|
||||
w_idx = spatial_idx % num_w_tiles
|
||||
|
||||
# Calculate positions
|
||||
t_start = t_idx * tile_latent_stride_num_frames
|
||||
h_start = h_idx * tile_latent_stride_height
|
||||
w_start = w_idx * tile_latent_stride_width
|
||||
|
||||
# Extract and process tile
|
||||
tile = z[:, :, t_start:t_start + tile_latent_min_num_frames + 1,
|
||||
h_start:h_start + tile_latent_min_height,
|
||||
w_start:w_start + tile_latent_min_width]
|
||||
|
||||
# Process tile
|
||||
tile = self._decode(tile)
|
||||
|
||||
if t_start > 0:
|
||||
tile = tile[:, :, 1:, :, :]
|
||||
|
||||
# Store metadata
|
||||
shape = tile.shape
|
||||
# Store decoded data (flattened)
|
||||
decoded_flat = tile.reshape(-1)
|
||||
local_results.append(decoded_flat)
|
||||
local_dim_metadata.append(shape)
|
||||
|
||||
results = torch.cat(local_results, dim=0).contiguous()
|
||||
del local_results
|
||||
torch.cuda.empty_cache()
|
||||
# first gather size to pad the results
|
||||
local_size = torch.tensor([results.size(0)],
|
||||
device=results.device,
|
||||
dtype=torch.int64)
|
||||
all_sizes = [
|
||||
torch.zeros(1, device=results.device, dtype=torch.int64)
|
||||
for _ in range(world_size)
|
||||
]
|
||||
dist.all_gather(all_sizes, local_size)
|
||||
max_size = max(size.item() for size in all_sizes)
|
||||
padded_results = torch.zeros(max_size, device=results.device)
|
||||
padded_results[:results.size(0)] = results
|
||||
del results
|
||||
torch.cuda.empty_cache()
|
||||
# Gather all results
|
||||
gathered_dim_metadata = [None] * world_size
|
||||
gathered_results = torch.zeros_like(padded_results).repeat(
|
||||
world_size, *[1] * len(padded_results.shape)
|
||||
).contiguous() # use contiguous to make sure it won't copy data in the following operations
|
||||
# TODO (PY): use fastvideo distributed methods
|
||||
dist.all_gather_into_tensor(gathered_results, padded_results)
|
||||
dist.all_gather_object(gathered_dim_metadata, local_dim_metadata)
|
||||
# Process gathered results
|
||||
data = [[[[] for _ in range(num_w_tiles)] for _ in range(num_h_tiles)]
|
||||
for _ in range(num_t_tiles)]
|
||||
for current_data, global_idx in self._parallel_data_generator(
|
||||
gathered_results, gathered_dim_metadata):
|
||||
t_idx = global_idx // total_spatial_tiles
|
||||
spatial_idx = global_idx % total_spatial_tiles
|
||||
h_idx = spatial_idx // num_w_tiles
|
||||
w_idx = spatial_idx % num_w_tiles
|
||||
data[t_idx][h_idx][w_idx] = current_data
|
||||
# Merge results
|
||||
result_slices = []
|
||||
last_slice_data = None
|
||||
for i, tem_data in enumerate(data):
|
||||
slice_data = self._merge_spatial_tiles(tem_data, blend_height, blend_width,
|
||||
self.tile_sample_stride_height, self.tile_sample_stride_width)
|
||||
if i > 0:
|
||||
slice_data = self.blend_t(last_slice_data, slice_data, blend_num_frames)
|
||||
result_slices.append(slice_data[:, :, :self.tile_sample_stride_num_frames, :, :])
|
||||
else:
|
||||
result_slices.append(slice_data[:, :, :self.tile_sample_stride_num_frames + 1, :, :])
|
||||
last_slice_data = slice_data
|
||||
dec = torch.cat(result_slices, dim=2)
|
||||
|
||||
return dec
|
||||
|
||||
def _merge_spatial_tiles(self, tiles, blend_height, blend_width, stride_height, stride_width):
|
||||
"""Helper function to merge spatial tiles with blending"""
|
||||
result_rows = []
|
||||
for i, row in enumerate(tiles):
|
||||
result_row = []
|
||||
for j, tile in enumerate(row):
|
||||
if i > 0:
|
||||
tile = self.blend_v(tiles[i - 1][j], tile,
|
||||
blend_height)
|
||||
if j > 0:
|
||||
tile = self.blend_h(row[j - 1], tile, blend_width)
|
||||
result_row.append(tile[:, :, :, :stride_height, :stride_width])
|
||||
result_rows.append(torch.cat(result_row, dim=-1))
|
||||
return torch.cat(result_rows, dim=-2)
|
||||
|
||||
def spatial_tiled_decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
r"""
|
||||
Decode a batch of images using a tiled decoder.
|
||||
|
||||
Args:
|
||||
z (`torch.Tensor`): Input batch of latent vectors.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
The decoded images.
|
||||
"""
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = z.shape
|
||||
sample_height = height * self.spatial_compression_ratio
|
||||
sample_width = width * self.spatial_compression_ratio
|
||||
|
||||
tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio
|
||||
tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio
|
||||
tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio
|
||||
tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio
|
||||
|
||||
blend_height = self.tile_sample_min_height - self.tile_sample_stride_height
|
||||
blend_width = self.tile_sample_min_width - self.tile_sample_stride_width
|
||||
|
||||
# Split z into overlapping tiles and decode them separately.
|
||||
# The tiles have an overlap to avoid seams between tiles.
|
||||
rows = []
|
||||
for i in range(0, height, tile_latent_stride_height):
|
||||
row = []
|
||||
for j in range(0, width, tile_latent_stride_width):
|
||||
tile = z[:, :, :, i : i + tile_latent_min_height, j : j + tile_latent_min_width]
|
||||
decoded = self._decode(tile)
|
||||
row.append(decoded)
|
||||
rows.append(row)
|
||||
return self._merge_spatial_tiles(rows, blend_height, blend_width,self.tile_sample_stride_height, self.tile_sample_stride_width)
|
||||
|
||||
def tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, num_channels, num_frames, height, width = x.shape
|
||||
latent_num_frames = (num_frames - 1) // self.temporal_compression_ratio + 1
|
||||
|
||||
tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio
|
||||
tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio
|
||||
blend_num_frames = tile_latent_min_num_frames - tile_latent_stride_num_frames
|
||||
|
||||
row = []
|
||||
for i in range(0, num_frames, self.tile_sample_stride_num_frames):
|
||||
tile = x[:, :, i : i + self.tile_sample_min_num_frames + 1, :, :]
|
||||
if self.use_tiling and (height > self.tile_sample_min_height or width > self.tile_sample_min_width):
|
||||
tile = self.spatial_tiled_encode(tile)
|
||||
else:
|
||||
tile = self._encode(tile)
|
||||
if i > 0:
|
||||
tile = tile[:, :, 1:, :, :]
|
||||
row.append(tile)
|
||||
result_row = []
|
||||
for i, tile in enumerate(row):
|
||||
if i > 0:
|
||||
tile = self.blend_t(row[i - 1], tile, blend_num_frames)
|
||||
result_row.append(tile[:, :, :tile_latent_stride_num_frames, :, :])
|
||||
else:
|
||||
result_row.append(tile[:, :, : tile_latent_stride_num_frames + 1, :, :])
|
||||
|
||||
enc = torch.cat(result_row, dim=2)[:, :, :latent_num_frames]
|
||||
return enc
|
||||
|
||||
def tiled_decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, num_channels, num_frames, height, width = z.shape
|
||||
num_sample_frames = (num_frames - 1) * self.temporal_compression_ratio + 1
|
||||
|
||||
tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio
|
||||
tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio
|
||||
tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio
|
||||
tile_latent_stride_num_frames = self.tile_sample_stride_num_frames // self.temporal_compression_ratio
|
||||
blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
|
||||
|
||||
row = []
|
||||
for i in range(0, num_frames, tile_latent_stride_num_frames):
|
||||
tile = z[:, :, i : i + tile_latent_min_num_frames + 1, :, :]
|
||||
if self.use_tiling and (tile.shape[-1] > tile_latent_min_width or tile.shape[-2] > tile_latent_min_height):
|
||||
decoded = self.spatial_tiled_decode(tile)
|
||||
else:
|
||||
decoded = self._decode(tile)
|
||||
if i > 0:
|
||||
decoded = decoded[:, :, 1:, :, :]
|
||||
row.append(decoded)
|
||||
result_row = []
|
||||
result_row = []
|
||||
for i, tile in enumerate(row):
|
||||
if i > 0:
|
||||
tile = self.blend_t(row[i - 1], tile, blend_num_frames)
|
||||
result_row.append(tile[:, :, : self.tile_sample_stride_num_frames, :, :])
|
||||
else:
|
||||
result_row.append(tile[:, :, : self.tile_sample_stride_num_frames + 1, :, :])
|
||||
|
||||
dec = torch.cat(result_row, dim=2)[:, :, :num_sample_frames]
|
||||
return dec
|
||||
|
||||
|
||||
|
||||
def enable_tiling(
|
||||
self,
|
||||
tile_sample_min_height: Optional[int] = None,
|
||||
tile_sample_min_width: Optional[int] = None,
|
||||
tile_sample_min_num_frames: Optional[int] = None,
|
||||
tile_sample_stride_height: Optional[float] = None,
|
||||
tile_sample_stride_width: Optional[float] = None,
|
||||
tile_sample_stride_num_frames: Optional[float] = None,
|
||||
) -> None:
|
||||
r"""
|
||||
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
|
||||
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
|
||||
processing larger images.
|
||||
|
||||
Args:
|
||||
tile_sample_min_height (`int`, *optional*):
|
||||
The minimum height required for a sample to be separated into tiles across the height dimension.
|
||||
tile_sample_min_width (`int`, *optional*):
|
||||
The minimum width required for a sample to be separated into tiles across the width dimension.
|
||||
tile_sample_min_num_frames (`int`, *optional*):
|
||||
The minimum number of frames required for a sample to be separated into tiles across the frame
|
||||
dimension.
|
||||
tile_sample_stride_height (`int`, *optional*):
|
||||
The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are
|
||||
no tiling artifacts produced across the height dimension.
|
||||
tile_sample_stride_width (`int`, *optional*):
|
||||
The stride between two consecutive horizontal tiles. This is to ensure that there are no tiling
|
||||
artifacts produced across the width dimension.
|
||||
tile_sample_stride_num_frames (`int`, *optional*):
|
||||
The stride between two consecutive frame tiles. This is to ensure that there are no tiling artifacts
|
||||
produced across the frame dimension.
|
||||
"""
|
||||
self.use_tiling = True
|
||||
self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height
|
||||
self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width
|
||||
self.tile_sample_min_num_frames = tile_sample_min_num_frames or self.tile_sample_min_num_frames
|
||||
self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height
|
||||
self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width
|
||||
self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames
|
||||
|
||||
|
||||
|
||||
def disable_tiling(self) -> None:
|
||||
r"""
|
||||
Disable tiled VAE decoding. If `enable_tiling` was previously enabled, this method will go back to computing
|
||||
decoding in one step.
|
||||
"""
|
||||
self.use_tiling = False
|
||||
|
||||
|
||||
|
||||
# adapted from https://github.com/huggingface/diffusers/blob/e7ffeae0a191f710881d1fbde00cd6ff025e81f2/src/diffusers/models/autoencoders/vae.py#L691
|
||||
class DiagonalGaussianDistribution(object):
|
||||
def __init__(self, parameters: torch.Tensor, deterministic: bool = False):
|
||||
self.parameters = parameters
|
||||
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
|
||||
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
||||
self.deterministic = deterministic
|
||||
self.std = torch.exp(0.5 * self.logvar)
|
||||
self.var = torch.exp(self.logvar)
|
||||
if self.deterministic:
|
||||
self.var = self.std = torch.zeros_like(
|
||||
self.mean, device=self.parameters.device, dtype=self.parameters.dtype
|
||||
)
|
||||
|
||||
def sample(self, generator: Optional[torch.Generator] = None) -> torch.Tensor:
|
||||
# make sure sample is on the same device as the parameters and has same dtype
|
||||
sample = torch.randn(
|
||||
self.mean.shape,
|
||||
generator=generator,
|
||||
device=self.parameters.device,
|
||||
dtype=self.parameters.dtype,
|
||||
)
|
||||
x = self.mean + self.std * sample
|
||||
return x
|
||||
|
||||
def kl(self, other: "DiagonalGaussianDistribution" = None) -> torch.Tensor:
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.0])
|
||||
else:
|
||||
if other is None:
|
||||
return 0.5 * torch.sum(
|
||||
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar,
|
||||
dim=[1, 2, 3],
|
||||
)
|
||||
else:
|
||||
return 0.5 * torch.sum(
|
||||
torch.pow(self.mean - other.mean, 2) / other.var
|
||||
+ self.var / other.var
|
||||
- 1.0
|
||||
- self.logvar
|
||||
+ other.logvar,
|
||||
dim=[1, 2, 3],
|
||||
)
|
||||
|
||||
def nll(self, sample: torch.Tensor, dims: Tuple[int, ...] = [1, 2, 3]) -> torch.Tensor:
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.0])
|
||||
logtwopi = np.log(2.0 * np.pi)
|
||||
return 0.5 * torch.sum(
|
||||
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
|
||||
dim=dims,
|
||||
)
|
||||
|
||||
def mode(self) -> torch.Tensor:
|
||||
return self.mean
|
||||
@@ -0,0 +1,799 @@
|
||||
# Copyright 2024 The Hunyuan Team, The HuggingFace Team and The FastVideo Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.checkpoint
|
||||
|
||||
from fastvideo.v1.models.vaes.common import DiagonalGaussianDistribution, ParallelTiledVAE
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
|
||||
from fastvideo.v1.models.utils import auto_attributes
|
||||
|
||||
def prepare_causal_attention_mask(
|
||||
num_frames: int, height_width: int, dtype: torch.dtype, device: torch.device, batch_size: int = None
|
||||
) -> torch.Tensor:
|
||||
indices = torch.arange(1, num_frames + 1, dtype=torch.int32, device=device)
|
||||
indices_blocks = indices.repeat_interleave(height_width)
|
||||
x, y = torch.meshgrid(indices_blocks, indices_blocks, indexing="xy")
|
||||
mask = torch.where(x <= y, 0, -float("inf")).to(dtype=dtype)
|
||||
|
||||
if batch_size is not None:
|
||||
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
return mask
|
||||
|
||||
class HunyuanVAEAttention(nn.Module):
|
||||
def __init__(self, in_channels, heads, dim_head, eps, norm_num_groups, bias):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
self.eps = eps
|
||||
self.norm_num_groups = norm_num_groups
|
||||
self.bias = bias
|
||||
|
||||
inner_dim = heads * dim_head
|
||||
|
||||
# Define the projection layers
|
||||
self.to_q = nn.Linear(in_channels, inner_dim, bias=bias)
|
||||
self.to_k = nn.Linear(in_channels, inner_dim, bias=bias)
|
||||
self.to_v = nn.Linear(in_channels, inner_dim, bias=bias)
|
||||
self.to_out = nn.Sequential(
|
||||
nn.Linear(inner_dim, in_channels, bias=bias)
|
||||
)
|
||||
|
||||
# Optional normalization layers
|
||||
self.group_norm = nn.GroupNorm(norm_num_groups, in_channels, eps=eps, affine=True)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
residual = hidden_states
|
||||
|
||||
batch_size, sequence_length, _ = hidden_states.shape
|
||||
|
||||
hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
# Project to query, key, value
|
||||
query = self.to_q(hidden_states)
|
||||
key = self.to_k(hidden_states)
|
||||
value = self.to_v(hidden_states)
|
||||
|
||||
# Reshape for multi-head attention
|
||||
head_dim = self.dim_head
|
||||
|
||||
query = query.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
|
||||
key = key.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# Perform scaled dot-product attention
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
# Reshape back
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# Linear projection
|
||||
hidden_states = self.to_out(hidden_states)
|
||||
|
||||
# Residual connection and rescale
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
return hidden_states
|
||||
|
||||
class HunyuanVideoCausalConv3d(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
kernel_size: Union[int, Tuple[int, int, int]] = 3,
|
||||
stride: Union[int, Tuple[int, int, int]] = 1,
|
||||
padding: Union[int, Tuple[int, int, int]] = 0,
|
||||
dilation: Union[int, Tuple[int, int, int]] = 1,
|
||||
bias: bool = True,
|
||||
pad_mode: str = "replicate",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
kernel_size = (kernel_size, kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size
|
||||
|
||||
self.pad_mode = pad_mode
|
||||
self.time_causal_padding = (
|
||||
kernel_size[0] // 2,
|
||||
kernel_size[0] // 2,
|
||||
kernel_size[1] // 2,
|
||||
kernel_size[1] // 2,
|
||||
kernel_size[2] - 1,
|
||||
0,
|
||||
)
|
||||
|
||||
self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, dilation, bias=bias)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = F.pad(hidden_states, self.time_causal_padding, mode=self.pad_mode)
|
||||
return self.conv(hidden_states)
|
||||
|
||||
|
||||
class HunyuanVideoUpsampleCausal3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: Optional[int] = None,
|
||||
kernel_size: int = 3,
|
||||
stride: int = 1,
|
||||
bias: bool = True,
|
||||
upsample_factor: Tuple[float, float, float] = (2, 2, 2),
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
out_channels = out_channels or in_channels
|
||||
self.upsample_factor = upsample_factor
|
||||
|
||||
self.conv = HunyuanVideoCausalConv3d(in_channels, out_channels, kernel_size, stride, bias=bias)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
num_frames = hidden_states.size(2)
|
||||
|
||||
first_frame, other_frames = hidden_states.split((1, num_frames - 1), dim=2)
|
||||
first_frame = F.interpolate(
|
||||
first_frame.squeeze(2), scale_factor=self.upsample_factor[1:], mode="nearest"
|
||||
).unsqueeze(2)
|
||||
|
||||
if num_frames > 1:
|
||||
# See: https://github.com/pytorch/pytorch/issues/81665
|
||||
# Unless you have a version of pytorch where non-contiguous implementation of F.interpolate
|
||||
# is fixed, this will raise either a runtime error, or fail silently with bad outputs.
|
||||
# If you are encountering an error here, make sure to try running encoding/decoding with
|
||||
# `vae.enable_tiling()` first. If that doesn't work, open an issue at:
|
||||
# https://github.com/huggingface/diffusers/issues
|
||||
other_frames = other_frames.contiguous()
|
||||
other_frames = F.interpolate(other_frames, scale_factor=self.upsample_factor, mode="nearest")
|
||||
hidden_states = torch.cat((first_frame, other_frames), dim=2)
|
||||
else:
|
||||
hidden_states = first_frame
|
||||
|
||||
hidden_states = self.conv(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoDownsampleCausal3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
channels: int,
|
||||
out_channels: Optional[int] = None,
|
||||
padding: int = 1,
|
||||
kernel_size: int = 3,
|
||||
bias: bool = True,
|
||||
stride=2,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
out_channels = out_channels or channels
|
||||
|
||||
self.conv = HunyuanVideoCausalConv3d(channels, out_channels, kernel_size, stride, padding, bias=bias)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.conv(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoResnetBlockCausal3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: Optional[int] = None,
|
||||
dropout: float = 0.0,
|
||||
groups: int = 32,
|
||||
eps: float = 1e-6,
|
||||
non_linearity: str = "silu",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
out_channels = out_channels or in_channels
|
||||
|
||||
self.nonlinearity = get_act_fn(non_linearity)
|
||||
|
||||
self.norm1 = nn.GroupNorm(groups, in_channels, eps=eps, affine=True)
|
||||
self.conv1 = HunyuanVideoCausalConv3d(in_channels, out_channels, 3, 1, 0)
|
||||
|
||||
self.norm2 = nn.GroupNorm(groups, out_channels, eps=eps, affine=True)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.conv2 = HunyuanVideoCausalConv3d(out_channels, out_channels, 3, 1, 0)
|
||||
|
||||
self.conv_shortcut = None
|
||||
if in_channels != out_channels:
|
||||
self.conv_shortcut = HunyuanVideoCausalConv3d(in_channels, out_channels, 1, 1, 0)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = hidden_states.contiguous()
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.norm1(hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
hidden_states = self.conv1(hidden_states)
|
||||
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
hidden_states = self.dropout(hidden_states)
|
||||
hidden_states = self.conv2(hidden_states)
|
||||
|
||||
if self.conv_shortcut is not None:
|
||||
residual = self.conv_shortcut(residual)
|
||||
|
||||
hidden_states = hidden_states + residual
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoMidBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_act_fn: str = "silu",
|
||||
resnet_groups: int = 32,
|
||||
add_attention: bool = True,
|
||||
attention_head_dim: int = 1,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)
|
||||
self.add_attention = add_attention
|
||||
|
||||
# There is always at least one resnet
|
||||
resnets = [
|
||||
HunyuanVideoResnetBlockCausal3D(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
non_linearity=resnet_act_fn,
|
||||
)
|
||||
]
|
||||
attentions = []
|
||||
|
||||
for _ in range(num_layers):
|
||||
if self.add_attention:
|
||||
attentions.append(
|
||||
HunyuanVAEAttention(
|
||||
in_channels,
|
||||
heads=in_channels // attention_head_dim,
|
||||
dim_head=attention_head_dim,
|
||||
eps=resnet_eps,
|
||||
norm_num_groups=resnet_groups,
|
||||
bias=True,
|
||||
)
|
||||
)
|
||||
else:
|
||||
attentions.append(None)
|
||||
|
||||
resnets.append(
|
||||
HunyuanVideoResnetBlockCausal3D(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
non_linearity=resnet_act_fn,
|
||||
)
|
||||
)
|
||||
|
||||
self.attentions = nn.ModuleList(attentions)
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
hidden_states = self._gradient_checkpointing_func(self.resnets[0], hidden_states)
|
||||
|
||||
for attn, resnet in zip(self.attentions, self.resnets[1:]):
|
||||
if attn is not None:
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 4, 1).flatten(1, 3)
|
||||
attention_mask = prepare_causal_attention_mask(
|
||||
num_frames, height * width, hidden_states.dtype, hidden_states.device, batch_size=batch_size
|
||||
)
|
||||
hidden_states = attn(hidden_states, attention_mask=attention_mask)
|
||||
hidden_states = hidden_states.unflatten(1, (num_frames, height, width)).permute(0, 4, 1, 2, 3)
|
||||
|
||||
hidden_states = self._gradient_checkpointing_func(resnet, hidden_states)
|
||||
|
||||
else:
|
||||
hidden_states = self.resnets[0](hidden_states)
|
||||
|
||||
for attn, resnet in zip(self.attentions, self.resnets[1:]):
|
||||
if attn is not None:
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 4, 1).flatten(1, 3)
|
||||
attention_mask = prepare_causal_attention_mask(
|
||||
num_frames, height * width, hidden_states.dtype, hidden_states.device, batch_size=batch_size
|
||||
)
|
||||
hidden_states = attn(hidden_states, attention_mask=attention_mask)
|
||||
hidden_states = hidden_states.unflatten(1, (num_frames, height, width)).permute(0, 4, 1, 2, 3)
|
||||
|
||||
hidden_states = resnet(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoDownBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_act_fn: str = "silu",
|
||||
resnet_groups: int = 32,
|
||||
add_downsample: bool = True,
|
||||
downsample_stride: int = 2,
|
||||
downsample_padding: int = 1,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
resnets = []
|
||||
|
||||
for i in range(num_layers):
|
||||
in_channels = in_channels if i == 0 else out_channels
|
||||
resnets.append(
|
||||
HunyuanVideoResnetBlockCausal3D(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
non_linearity=resnet_act_fn,
|
||||
)
|
||||
)
|
||||
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
if add_downsample:
|
||||
self.downsamplers = nn.ModuleList(
|
||||
[
|
||||
HunyuanVideoDownsampleCausal3D(
|
||||
out_channels,
|
||||
out_channels=out_channels,
|
||||
padding=downsample_padding,
|
||||
stride=downsample_stride,
|
||||
)
|
||||
]
|
||||
)
|
||||
else:
|
||||
self.downsamplers = None
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = self._gradient_checkpointing_func(resnet, hidden_states)
|
||||
else:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = resnet(hidden_states)
|
||||
|
||||
if self.downsamplers is not None:
|
||||
for downsampler in self.downsamplers:
|
||||
hidden_states = downsampler(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoUpBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_act_fn: str = "silu",
|
||||
resnet_groups: int = 32,
|
||||
add_upsample: bool = True,
|
||||
upsample_scale_factor: Tuple[int, int, int] = (2, 2, 2),
|
||||
) -> None:
|
||||
super().__init__()
|
||||
resnets = []
|
||||
|
||||
for i in range(num_layers):
|
||||
input_channels = in_channels if i == 0 else out_channels
|
||||
|
||||
resnets.append(
|
||||
HunyuanVideoResnetBlockCausal3D(
|
||||
in_channels=input_channels,
|
||||
out_channels=out_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
non_linearity=resnet_act_fn,
|
||||
)
|
||||
)
|
||||
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
if add_upsample:
|
||||
self.upsamplers = nn.ModuleList(
|
||||
[
|
||||
HunyuanVideoUpsampleCausal3D(
|
||||
out_channels,
|
||||
out_channels=out_channels,
|
||||
upsample_factor=upsample_scale_factor,
|
||||
)
|
||||
]
|
||||
)
|
||||
else:
|
||||
self.upsamplers = None
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = self._gradient_checkpointing_func(resnet, hidden_states)
|
||||
|
||||
else:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = resnet(hidden_states)
|
||||
|
||||
if self.upsamplers is not None:
|
||||
for upsampler in self.upsamplers:
|
||||
hidden_states = upsampler(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoEncoder3D(nn.Module):
|
||||
r"""
|
||||
Causal encoder for 3D video-like data introduced in [Hunyuan Video](https://huggingface.co/papers/2412.03603).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
down_block_types: Tuple[str, ...] = (
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
),
|
||||
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
|
||||
layers_per_block: int = 2,
|
||||
norm_num_groups: int = 32,
|
||||
act_fn: str = "silu",
|
||||
double_z: bool = True,
|
||||
mid_block_add_attention=True,
|
||||
temporal_compression_ratio: int = 4,
|
||||
spatial_compression_ratio: int = 8,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.conv_in = HunyuanVideoCausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
|
||||
self.mid_block = None
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
|
||||
output_channel = block_out_channels[0]
|
||||
for i, down_block_type in enumerate(down_block_types):
|
||||
if down_block_type != "HunyuanVideoDownBlock3D":
|
||||
raise ValueError(f"Unsupported down_block_type: {down_block_type}")
|
||||
|
||||
input_channel = output_channel
|
||||
output_channel = block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
num_spatial_downsample_layers = int(np.log2(spatial_compression_ratio))
|
||||
num_time_downsample_layers = int(np.log2(temporal_compression_ratio))
|
||||
|
||||
if temporal_compression_ratio == 4:
|
||||
add_spatial_downsample = bool(i < num_spatial_downsample_layers)
|
||||
add_time_downsample = bool(
|
||||
i >= (len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block
|
||||
)
|
||||
elif temporal_compression_ratio == 8:
|
||||
add_spatial_downsample = bool(i < num_spatial_downsample_layers)
|
||||
add_time_downsample = bool(i < num_time_downsample_layers)
|
||||
else:
|
||||
raise ValueError(f"Unsupported time_compression_ratio: {temporal_compression_ratio}")
|
||||
|
||||
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
|
||||
downsample_stride_T = (2,) if add_time_downsample else (1,)
|
||||
downsample_stride = tuple(downsample_stride_T + downsample_stride_HW)
|
||||
|
||||
down_block = HunyuanVideoDownBlock3D(
|
||||
num_layers=layers_per_block,
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
add_downsample=bool(add_spatial_downsample or add_time_downsample),
|
||||
resnet_eps=1e-6,
|
||||
resnet_act_fn=act_fn,
|
||||
resnet_groups=norm_num_groups,
|
||||
downsample_stride=downsample_stride,
|
||||
downsample_padding=0,
|
||||
)
|
||||
|
||||
self.down_blocks.append(down_block)
|
||||
|
||||
self.mid_block = HunyuanVideoMidBlock3D(
|
||||
in_channels=block_out_channels[-1],
|
||||
resnet_eps=1e-6,
|
||||
resnet_act_fn=act_fn,
|
||||
attention_head_dim=block_out_channels[-1],
|
||||
resnet_groups=norm_num_groups,
|
||||
add_attention=mid_block_add_attention,
|
||||
)
|
||||
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6)
|
||||
self.conv_act = nn.SiLU()
|
||||
|
||||
conv_out_channels = 2 * out_channels if double_z else out_channels
|
||||
self.conv_out = HunyuanVideoCausalConv3d(block_out_channels[-1], conv_out_channels, kernel_size=3)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.conv_in(hidden_states)
|
||||
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for down_block in self.down_blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(down_block, hidden_states)
|
||||
|
||||
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
|
||||
else:
|
||||
for down_block in self.down_blocks:
|
||||
hidden_states = down_block(hidden_states)
|
||||
|
||||
hidden_states = self.mid_block(hidden_states)
|
||||
|
||||
hidden_states = self.conv_norm_out(hidden_states)
|
||||
hidden_states = self.conv_act(hidden_states)
|
||||
hidden_states = self.conv_out(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoDecoder3D(nn.Module):
|
||||
r"""
|
||||
Causal decoder for 3D video-like data introduced in [Hunyuan Video](https://huggingface.co/papers/2412.03603).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
up_block_types: Tuple[str, ...] = (
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
),
|
||||
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
|
||||
layers_per_block: int = 2,
|
||||
norm_num_groups: int = 32,
|
||||
act_fn: str = "silu",
|
||||
mid_block_add_attention=True,
|
||||
time_compression_ratio: int = 4,
|
||||
spatial_compression_ratio: int = 8,
|
||||
):
|
||||
super().__init__()
|
||||
self.layers_per_block = layers_per_block
|
||||
|
||||
self.conv_in = HunyuanVideoCausalConv3d(in_channels, block_out_channels[-1], kernel_size=3, stride=1)
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
|
||||
# mid
|
||||
self.mid_block = HunyuanVideoMidBlock3D(
|
||||
in_channels=block_out_channels[-1],
|
||||
resnet_eps=1e-6,
|
||||
resnet_act_fn=act_fn,
|
||||
attention_head_dim=block_out_channels[-1],
|
||||
resnet_groups=norm_num_groups,
|
||||
add_attention=mid_block_add_attention,
|
||||
)
|
||||
|
||||
# up
|
||||
reversed_block_out_channels = list(reversed(block_out_channels))
|
||||
output_channel = reversed_block_out_channels[0]
|
||||
for i, up_block_type in enumerate(up_block_types):
|
||||
if up_block_type != "HunyuanVideoUpBlock3D":
|
||||
raise ValueError(f"Unsupported up_block_type: {up_block_type}")
|
||||
|
||||
prev_output_channel = output_channel
|
||||
output_channel = reversed_block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
num_spatial_upsample_layers = int(np.log2(spatial_compression_ratio))
|
||||
num_time_upsample_layers = int(np.log2(time_compression_ratio))
|
||||
|
||||
if time_compression_ratio == 4:
|
||||
add_spatial_upsample = bool(i < num_spatial_upsample_layers)
|
||||
add_time_upsample = bool(
|
||||
i >= len(block_out_channels) - 1 - num_time_upsample_layers and not is_final_block
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}")
|
||||
|
||||
upsample_scale_factor_HW = (2, 2) if add_spatial_upsample else (1, 1)
|
||||
upsample_scale_factor_T = (2,) if add_time_upsample else (1,)
|
||||
upsample_scale_factor = tuple(upsample_scale_factor_T + upsample_scale_factor_HW)
|
||||
|
||||
up_block = HunyuanVideoUpBlock3D(
|
||||
num_layers=self.layers_per_block + 1,
|
||||
in_channels=prev_output_channel,
|
||||
out_channels=output_channel,
|
||||
add_upsample=bool(add_spatial_upsample or add_time_upsample),
|
||||
upsample_scale_factor=upsample_scale_factor,
|
||||
resnet_eps=1e-6,
|
||||
resnet_act_fn=act_fn,
|
||||
resnet_groups=norm_num_groups,
|
||||
)
|
||||
|
||||
self.up_blocks.append(up_block)
|
||||
prev_output_channel = output_channel
|
||||
|
||||
# out
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6)
|
||||
self.conv_act = nn.SiLU()
|
||||
self.conv_out = HunyuanVideoCausalConv3d(block_out_channels[0], out_channels, kernel_size=3)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.conv_in(hidden_states)
|
||||
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
|
||||
|
||||
for up_block in self.up_blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(up_block, hidden_states)
|
||||
else:
|
||||
hidden_states = self.mid_block(hidden_states)
|
||||
|
||||
for up_block in self.up_blocks:
|
||||
hidden_states = up_block(hidden_states)
|
||||
|
||||
# post-process
|
||||
hidden_states = self.conv_norm_out(hidden_states)
|
||||
hidden_states = self.conv_act(hidden_states)
|
||||
hidden_states = self.conv_out(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
|
||||
r"""
|
||||
A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos.
|
||||
Introduced in [HunyuanVideo](https://huggingface.co/papers/2412.03603).
|
||||
|
||||
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
|
||||
for all models (such as downloading or saving).
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
@auto_attributes
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
latent_channels: int = 16,
|
||||
down_block_types: Tuple[str, ...] = (
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
),
|
||||
up_block_types: Tuple[str, ...] = (
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
),
|
||||
block_out_channels: Tuple[int] = (128, 256, 512, 512),
|
||||
layers_per_block: int = 2,
|
||||
act_fn: str = "silu",
|
||||
norm_num_groups: int = 32,
|
||||
scaling_factor: float = 0.476986,
|
||||
spatial_compression_ratio: int = 8,
|
||||
temporal_compression_ratio: int = 4,
|
||||
mid_block_add_attention: bool = True,
|
||||
load_encoder: bool = True,
|
||||
load_decoder: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
# TODO(will): only pass in config. We do this by manually defining a
|
||||
# config for hunyuan vae
|
||||
self.block_out_channels = block_out_channels
|
||||
|
||||
if load_encoder:
|
||||
self.encoder = HunyuanVideoEncoder3D(
|
||||
in_channels=in_channels,
|
||||
out_channels=latent_channels,
|
||||
down_block_types=down_block_types,
|
||||
block_out_channels=block_out_channels,
|
||||
layers_per_block=layers_per_block,
|
||||
norm_num_groups=norm_num_groups,
|
||||
act_fn=act_fn,
|
||||
double_z=True,
|
||||
mid_block_add_attention=mid_block_add_attention,
|
||||
temporal_compression_ratio=temporal_compression_ratio,
|
||||
spatial_compression_ratio=spatial_compression_ratio,
|
||||
)
|
||||
self.quant_conv = nn.Conv3d(2 * latent_channels, 2 * latent_channels, kernel_size=1)
|
||||
|
||||
if load_decoder:
|
||||
self.decoder = HunyuanVideoDecoder3D(
|
||||
in_channels=latent_channels,
|
||||
out_channels=out_channels,
|
||||
up_block_types=up_block_types,
|
||||
block_out_channels=block_out_channels,
|
||||
layers_per_block=layers_per_block,
|
||||
norm_num_groups=norm_num_groups,
|
||||
act_fn=act_fn,
|
||||
time_compression_ratio=temporal_compression_ratio,
|
||||
spatial_compression_ratio=spatial_compression_ratio,
|
||||
mid_block_add_attention=mid_block_add_attention,
|
||||
)
|
||||
self.post_quant_conv = nn.Conv3d(latent_channels, latent_channels, kernel_size=1)
|
||||
|
||||
|
||||
|
||||
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
|
||||
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
|
||||
# intermediate tiles together, the memory requirement can be lowered.
|
||||
self.use_tiling = True
|
||||
|
||||
# The minimal tile height and width for spatial tiling to be used
|
||||
self.tile_sample_min_height = 256
|
||||
self.tile_sample_min_width = 256
|
||||
self.tile_sample_min_num_frames = 16
|
||||
|
||||
# The minimal distance between two spatial tiles
|
||||
self.tile_sample_stride_height = 192
|
||||
self.tile_sample_stride_width = 192
|
||||
self.tile_sample_stride_num_frames = 12
|
||||
|
||||
|
||||
def _encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.encoder(x)
|
||||
enc = self.quant_conv(x)
|
||||
return enc
|
||||
|
||||
def _decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
z = self.post_quant_conv(z)
|
||||
dec = self.decoder(z)
|
||||
return dec
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.Tensor,
|
||||
sample_posterior: bool = False,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
Args:
|
||||
sample (`torch.Tensor`): Input sample.
|
||||
sample_posterior (`bool`, *optional*, defaults to `False`):
|
||||
Whether to sample from the posterior.
|
||||
"""
|
||||
x = sample
|
||||
posterior = self.encode(x).latent_dist
|
||||
if sample_posterior:
|
||||
z = posterior.sample(generator=generator)
|
||||
else:
|
||||
z = posterior.mode()
|
||||
dec = self.decode(z)
|
||||
return dec
|
||||
|
||||
@@ -0,0 +1,260 @@
|
||||
"""
|
||||
Diffusion pipelines for fastvideo.v1.
|
||||
|
||||
This package contains diffusion pipelines for generating videos and images.
|
||||
"""
|
||||
import os
|
||||
import json
|
||||
from copy import deepcopy
|
||||
|
||||
from typing import Dict, Optional, Type, Any
|
||||
|
||||
from fastvideo.v1.pipelines.pipeline_registry import PipelineRegistry
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from huggingface_hub import snapshot_download
|
||||
from transformers import PretrainedConfig
|
||||
from fastvideo.v1.models.hf_transformer_utils import get_hf_config, get_diffusers_config
|
||||
from fastvideo.v1.models import get_scheduler
|
||||
import glob
|
||||
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
|
||||
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
|
||||
|
||||
logger = init_logger(__name__)
|
||||
import tempfile
|
||||
import filelock
|
||||
import hashlib
|
||||
# Then import the base classes
|
||||
from fastvideo.v1.pipelines.composed.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
DiffusionPipelineOutput
|
||||
)
|
||||
|
||||
def get_pipeline_type(inference_args: InferenceArgs) -> str:
|
||||
# hardcode for now
|
||||
return "hunyuan_video"
|
||||
|
||||
|
||||
|
||||
def get_lock(model_name_or_path: str):
|
||||
lock_dir = tempfile.gettempdir()
|
||||
os.makedirs(os.path.dirname(lock_dir), exist_ok=True)
|
||||
model_name = model_name_or_path.replace("/", "-")
|
||||
hash_name = hashlib.sha256(model_name.encode()).hexdigest()
|
||||
# add hash to avoid conflict with old users' lock files
|
||||
lock_file_name = hash_name + model_name + ".lock"
|
||||
# mode 0o666 is required for the filelock to be shared across users
|
||||
lock = filelock.FileLock(os.path.join(lock_dir, lock_file_name),
|
||||
mode=0o666)
|
||||
return lock
|
||||
|
||||
|
||||
|
||||
def maybe_download_model(model_path: str) -> str:
|
||||
"""
|
||||
Check if the model path is a Hugging Face Hub model ID and download it if needed.
|
||||
|
||||
Args:
|
||||
model_path: Local path or Hugging Face Hub model ID
|
||||
|
||||
Returns:
|
||||
Local path to the model
|
||||
"""
|
||||
|
||||
# If the path exists locally, return it
|
||||
if os.path.exists(model_path):
|
||||
logger.info(f"Model already exists locally at {model_path}")
|
||||
return model_path
|
||||
|
||||
# Otherwise, assume it's a HF Hub model ID and try to download it
|
||||
try:
|
||||
logger.info(f"Downloading model snapshot from HF Hub for {model_path}...")
|
||||
with get_lock(model_path):
|
||||
local_path = snapshot_download(
|
||||
repo_id=model_path,
|
||||
ignore_patterns=["*.onnx", "*.msgpack"],
|
||||
)
|
||||
logger.info(f"Downloaded model to {local_path}")
|
||||
return local_path
|
||||
except Exception as e:
|
||||
raise ValueError(f"Could not find model at {model_path} and failed to download from HF Hub: {e}")
|
||||
|
||||
def verify_model_config_and_directory(model_path: str) -> dict:
|
||||
"""
|
||||
Verify that the model directory contains a valid diffusers configuration.
|
||||
|
||||
Args:
|
||||
model_path: Path to the model directory
|
||||
|
||||
Returns:
|
||||
The loaded model configuration as a dictionary
|
||||
"""
|
||||
|
||||
# Check for model_index.json which is required for diffusers models
|
||||
config_path = os.path.join(model_path, "model_index.json")
|
||||
if not os.path.exists(config_path):
|
||||
raise ValueError(
|
||||
f"Model directory {model_path} does not contain model_index.json. "
|
||||
"Only Hugging Face diffusers format is supported."
|
||||
)
|
||||
|
||||
# Check for transformer and vae directories
|
||||
transformer_dir = os.path.join(model_path, "transformer")
|
||||
vae_dir = os.path.join(model_path, "vae")
|
||||
|
||||
if not os.path.exists(transformer_dir):
|
||||
raise ValueError(f"Model directory {model_path} does not contain a transformer/ directory.")
|
||||
|
||||
if not os.path.exists(vae_dir):
|
||||
raise ValueError(f"Model directory {model_path} does not contain a vae/ directory.")
|
||||
|
||||
# Load the config
|
||||
try:
|
||||
with open(config_path, "r") as f:
|
||||
config = json.load(f)
|
||||
except Exception as e:
|
||||
raise ValueError(f"Failed to load model configuration from {config_path}: {e}")
|
||||
|
||||
# Verify diffusers version exists
|
||||
if "_diffusers_version" not in config:
|
||||
raise ValueError(f"model_index.json does not contain _diffusers_version")
|
||||
|
||||
logger.info(f"Diffusers version: {config['_diffusers_version']}")
|
||||
return config
|
||||
|
||||
def load_pipeline_module(module_name: str, component_model_path: str, transformers_or_diffusers: str, architecture: str, inference_args: InferenceArgs) -> Any:
|
||||
"""
|
||||
Load a pipeline module using the appropriate loader.
|
||||
|
||||
Args:
|
||||
module_name: Name of the module (e.g., "vae", "text_encoder", "transformer", "scheduler")
|
||||
component_model_path: Path to the component model
|
||||
transformers_or_diffusers: Whether the module is from transformers or diffusers
|
||||
architecture: Architecture of the component model
|
||||
inference_args: Inference arguments
|
||||
|
||||
Returns:
|
||||
The loaded module
|
||||
"""
|
||||
return PipelineComponentLoader.load_module(
|
||||
module_name=module_name,
|
||||
component_model_path=component_model_path,
|
||||
transformers_or_diffusers=transformers_or_diffusers,
|
||||
architecture=architecture,
|
||||
inference_args=inference_args
|
||||
)
|
||||
|
||||
def load_pipeline_modules(model_path: str, config: Dict, inference_args: InferenceArgs) -> dict[str, Any]:
|
||||
"""
|
||||
Load the pipeline modules from the config.
|
||||
|
||||
Args:
|
||||
config: The model_index.json config
|
||||
inference_args: Inference arguments
|
||||
|
||||
Returns:
|
||||
Dictionary mapping module names to loaded modules
|
||||
"""
|
||||
logger.info(f"Loading pipeline modules from config: {config}")
|
||||
modules_config = deepcopy(config)
|
||||
|
||||
# remove keys that are not pipeline modules
|
||||
modules_config.pop("_class_name")
|
||||
modules_config.pop("_diffusers_version")
|
||||
|
||||
# some sanity checks
|
||||
assert len(modules_config) > 1, "model_index.json must contain at least one pipeline module"
|
||||
|
||||
required_modules = ["vae", "text_encoder", "transformer", "scheduler", "tokenizer"]
|
||||
for module_name in required_modules:
|
||||
if module_name not in modules_config:
|
||||
raise ValueError(f"model_index.json must contain a {module_name} module")
|
||||
logger.info(f"Diffusers config passed sanity checks")
|
||||
|
||||
|
||||
# all the component models used by the pipeline
|
||||
pipeline_modules = {}
|
||||
for module_name, (transformers_or_diffusers, architecture) in modules_config.items():
|
||||
component_model_path = os.path.join(model_path, module_name)
|
||||
module = load_pipeline_module(
|
||||
module_name,
|
||||
component_model_path,
|
||||
transformers_or_diffusers,
|
||||
architecture,
|
||||
inference_args,
|
||||
)
|
||||
|
||||
pipeline_modules[module_name] = module
|
||||
|
||||
|
||||
# Check if all required modules were loaded
|
||||
for module_name in required_modules:
|
||||
if module_name not in pipeline_modules or pipeline_modules[module_name] is None:
|
||||
logger.warning(f"Required module {module_name} was not loaded properly")
|
||||
|
||||
return pipeline_modules
|
||||
|
||||
def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
|
||||
"""
|
||||
Only works with valid hf diffusers configs. (model_index.json)
|
||||
We want to build a pipeline based on the inference args mode_path:
|
||||
1. download the model from the hub if it's not already downloaded
|
||||
2. verify the model config and directory
|
||||
3. based on the config, determine the pipeline class
|
||||
4. parse the config to get the model components (vae, text_encoders, etc...)
|
||||
5. the pipeline loader class will use the model component names and paths to load
|
||||
6. the pipeline class will be composed of the models returned by the pipeline loader
|
||||
"""
|
||||
# Get pipeline type
|
||||
model_path = inference_args.model_path
|
||||
model_path = maybe_download_model(model_path)
|
||||
# inference_args.downloaded_model_path = model_path
|
||||
logger.info(f"Model path: {model_path}")
|
||||
config = verify_model_config_and_directory(model_path)
|
||||
|
||||
pipeline_architecture = config.get("_class_name")
|
||||
if pipeline_architecture is None:
|
||||
raise ValueError("Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported.")
|
||||
|
||||
pipeline_cls, pipeline_architecture = PipelineRegistry.resolve_pipeline_cls(pipeline_architecture)
|
||||
|
||||
# instantiate the pipeline
|
||||
pipeline = pipeline_cls()
|
||||
|
||||
|
||||
pipeline_modules = load_pipeline_modules(model_path, config, inference_args)
|
||||
|
||||
|
||||
logger.info(f"Initializing encoders")
|
||||
pipeline.initialize_encoders(pipeline_modules, inference_args)
|
||||
|
||||
logger.info(f"Registering modules")
|
||||
pipeline.register_modules(pipeline_modules)
|
||||
|
||||
logger.info(f"Setting up pipeline")
|
||||
pipeline.setup_pipeline(inference_args)
|
||||
|
||||
logger.info(f"Initializing pipeline")
|
||||
pipeline.initialize_pipeline(inference_args)
|
||||
|
||||
# pipeline is now initialized and ready to use
|
||||
return pipeline
|
||||
|
||||
|
||||
|
||||
def list_available_pipelines() -> Dict[str, Type[Any]]:
|
||||
"""
|
||||
List all available pipeline types.
|
||||
|
||||
Returns:
|
||||
A dictionary of pipeline names to pipeline classes.
|
||||
"""
|
||||
return PipelineRegistry.list()
|
||||
|
||||
__all__ = [
|
||||
"build_pipeline",
|
||||
"list_available_pipelines",
|
||||
"ComposedPipelineBase",
|
||||
"DiffusionPipelineOutput",
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
"""
|
||||
Composed pipelines for diffusion models.
|
||||
|
||||
This package contains pipelines that are composed of multiple stages.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.pipelines.composed.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
DiffusionPipelineOutput,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ComposedPipelineBase",
|
||||
"DiffusionPipelineOutput",
|
||||
]
|
||||
|
||||
# Note: Do not import TextToVideoPipeline here to avoid circular imports.
|
||||
# It will be imported in the main pipelines/__init__.py file.
|
||||
@@ -0,0 +1,147 @@
|
||||
"""
|
||||
Base class for composed pipelines.
|
||||
|
||||
This module defines the base class for pipelines that are composed of multiple stages.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Union, Tuple
|
||||
import torch
|
||||
import numpy as np
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.pipelines.stages import PipelineStage
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@dataclass
|
||||
class DiffusionPipelineOutput:
|
||||
"""Output from a diffusion pipeline."""
|
||||
videos: Union[torch.Tensor, np.ndarray]
|
||||
|
||||
|
||||
class ComposedPipelineBase(ABC):
|
||||
"""
|
||||
Base class for pipelines composed of multiple stages.
|
||||
|
||||
This class provides the framework for creating pipelines by composing multiple
|
||||
stages together. Each stage is responsible for a specific part of the diffusion
|
||||
process, and the pipeline orchestrates the execution of these stages.
|
||||
"""
|
||||
|
||||
is_video_pipeline: bool = False # To be overridden by video pipelines
|
||||
|
||||
def __init__(self):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
The pipeline should be completely stateless and not hold any batch
|
||||
state.
|
||||
"""
|
||||
self._stages: List[PipelineStage] = []
|
||||
self._modules: Dict[str, Any] = {}
|
||||
self._stage_name_mapping: Dict[str, PipelineStage] = {}
|
||||
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
"""Get the device for this pipeline."""
|
||||
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
@property
|
||||
def modules(self) -> Dict[str, Any]:
|
||||
"""Get all modules used by this pipeline."""
|
||||
return self._modules
|
||||
|
||||
@abstractmethod
|
||||
def setup_pipeline(self, inference_args: InferenceArgs):
|
||||
"""
|
||||
Setup the pipeline.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def initialize_encoders(self, modules: Dict[str, Any], inference_args: InferenceArgs):
|
||||
"""
|
||||
Initialize the encoders. Will remove the encoders/tokenizers modules from the
|
||||
modules. Will add the TextEncoder or ImageEncoder to the modules.
|
||||
"""
|
||||
...
|
||||
|
||||
def register_modules(self, modules: Dict[str, Any]):
|
||||
"""
|
||||
Register modules with the pipeline and its stages.
|
||||
|
||||
We will use the _module_name_mapping to map the module names used
|
||||
in the Diffusers config to the internal names (how it can be accessed
|
||||
in the pipeline).
|
||||
|
||||
Args:
|
||||
modules: The modules to register.
|
||||
"""
|
||||
self._modules.update(modules)
|
||||
# Register modules with self
|
||||
for name, module in modules.items():
|
||||
setattr(self, name, module)
|
||||
# self._modules[name] = module
|
||||
|
||||
# Register modules with stages that need them
|
||||
for stage in self._stages:
|
||||
stage.register_modules(modules)
|
||||
# TODO(will): perhaps we should not register all modules with the
|
||||
# stage. See below.
|
||||
|
||||
# stage_modules = {}
|
||||
# for name, module in mapped_modules.items():
|
||||
# if hasattr(stage, f"needs_{name}") and getattr(stage, f"needs_{name}"):
|
||||
# stage_modules[name] = module
|
||||
|
||||
# if stage_modules:
|
||||
# stage.register_modules(**stage_modules)
|
||||
|
||||
def add_stage(self, name: str, stage: PipelineStage):
|
||||
assert self._modules is not None, "No modules are registered"
|
||||
stage.register_modules(self._modules)
|
||||
self._stages.append(stage)
|
||||
self._stage_name_mapping[name] = stage
|
||||
setattr(self, name, stage)
|
||||
|
||||
|
||||
|
||||
|
||||
# TODO(will): don't hardcode no_grad
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> DiffusionPipelineOutput:
|
||||
"""
|
||||
Generate a video or image using the pipeline.
|
||||
|
||||
Args:
|
||||
prompt: The prompt(s) to guide generation.
|
||||
negative_prompt: The negative prompt(s) to guide generation.
|
||||
height: The height of the generated video/image.
|
||||
width: The width of the generated video/image.
|
||||
num_frames: The number of frames to generate (for video).
|
||||
num_inference_steps: The number of inference steps.
|
||||
guidance_scale: The scale for classifier-free guidance.
|
||||
num_videos_per_prompt: The number of videos to generate per prompt.
|
||||
generator: The random number generator.
|
||||
latents: The initial latents.
|
||||
output_type: The output type.
|
||||
**kwargs: Additional arguments.
|
||||
|
||||
Returns:
|
||||
The generated video or image.
|
||||
"""
|
||||
# Execute each stage
|
||||
for stage in self._stages:
|
||||
batch = stage(batch, inference_args)
|
||||
|
||||
# Return the output
|
||||
return DiffusionPipelineOutput(videos=batch.output)
|
||||
@@ -0,0 +1,13 @@
|
||||
"""
|
||||
HunYuan pipeline implementations.
|
||||
|
||||
This package contains implementations of diffusion pipelines for HunYuan models.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.pipelines.implementations.hunyuan.hunyuan_pipeline import (
|
||||
EntryClass,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"EntryClass",
|
||||
]
|
||||
@@ -0,0 +1,90 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
__all__ = [
|
||||
"PROMPT_TEMPLATE",
|
||||
"MODEL_BASE",
|
||||
"PRECISIONS",
|
||||
"NORMALIZATION_TYPE",
|
||||
"ACTIVATION_TYPE",
|
||||
"VAE_PATH",
|
||||
"TEXT_ENCODER_PATH",
|
||||
"TOKENIZER_PATH",
|
||||
"TEXT_PROJECTION",
|
||||
"DATA_TYPE",
|
||||
"NEGATIVE_PROMPT",
|
||||
]
|
||||
|
||||
PRECISION_TO_TYPE = {
|
||||
"fp32": torch.float32,
|
||||
"fp16": torch.float16,
|
||||
"bf16": torch.bfloat16,
|
||||
}
|
||||
|
||||
# =================== Constant Values =====================
|
||||
# Computation scale factor, 1P = 1_000_000_000_000_000. Tensorboard will display the value in PetaFLOPS to avoid
|
||||
# overflow error when tensorboard logging values.
|
||||
|
||||
# When using decoder-only models, we must provide a prompt template to instruct the text encoder
|
||||
# on how to generate the text.
|
||||
# --------------------------------------------------------------------
|
||||
|
||||
# TODO: Many models will have this model specific prompt template. Create a centralized place to store them.
|
||||
# TODO: add inference arg: --use-template to use the model specific prompt template. Else use no template
|
||||
PROMPT_TEMPLATE_ENCODE = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
|
||||
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
|
||||
PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
|
||||
"1. The main content and theme of the video."
|
||||
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
|
||||
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
|
||||
"4. background environment, light, style and atmosphere."
|
||||
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
|
||||
|
||||
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
|
||||
|
||||
PROMPT_TEMPLATE = {
|
||||
"image": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE,
|
||||
"crop_start": 36,
|
||||
},
|
||||
"video": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
|
||||
"crop_start": 95,
|
||||
},
|
||||
}
|
||||
|
||||
# ======================= Model ======================
|
||||
PRECISIONS = {"fp32", "fp16", "bf16"}
|
||||
NORMALIZATION_TYPE = {"layer", "rms"}
|
||||
ACTIVATION_TYPE = {"relu", "silu", "gelu", "gelu_tanh"}
|
||||
|
||||
# =================== Model Path =====================
|
||||
MODEL_BASE = os.getenv("MODEL_BASE", "./data/hunyuan")
|
||||
|
||||
# =================== Data =======================
|
||||
DATA_TYPE = {"image", "video", "image_video"}
|
||||
|
||||
# 3D VAE
|
||||
VAE_PATH = {"884-16c-hy": f"{MODEL_BASE}/hunyuan-video-t2v-720p/vae"}
|
||||
|
||||
# Text Encoder
|
||||
TEXT_ENCODER_PATH = {
|
||||
"clipL": f"{MODEL_BASE}/text_encoder_2",
|
||||
"llm": f"{MODEL_BASE}/text_encoder",
|
||||
}
|
||||
|
||||
# Tokenizer
|
||||
TOKENIZER_PATH = {
|
||||
"clipL": f"{MODEL_BASE}/text_encoder_2",
|
||||
"llm": f"{MODEL_BASE}/text_encoder",
|
||||
}
|
||||
|
||||
TEXT_PROJECTION = {
|
||||
"linear", # Default, an nn.Linear() layer
|
||||
"single_refiner", # Single TokenRefiner. Refer to LI-DiT
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
"""
|
||||
HunYuan video diffusion pipeline implementation.
|
||||
|
||||
This module contains an implementation of the HunYuan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from typing import Union, Any, Dict
|
||||
import torch
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
|
||||
from fastvideo.v1.pipelines.composed import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (
|
||||
InputValidationStage,
|
||||
PromptEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
LatentPreparationStage,
|
||||
ConditioningStage,
|
||||
DenoisingStage,
|
||||
DecodingStage,
|
||||
PostProcessingStage,
|
||||
)
|
||||
from fastvideo.v1.pipelines.stages.prompt_encoding import PromptEncodingStage
|
||||
# from fastvideo.v1.pipelines.stages.timestep_preparation import FlowMatchingTimestepPreparationStage
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
# from fastvideo.v1.pipelines.composed.composed_pipeline_base import DiffusionPipelineOutput
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
# TODO(will): move PRECISION_TO_TYPE to better place
|
||||
from .constants import PROMPT_TEMPLATE, PRECISION_TO_TYPE
|
||||
|
||||
from diffusers.utils import BaseOutput
|
||||
import numpy as np
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# class HunyuanLatentPreparationStage(LatentPreparationStage):
|
||||
# def _call_implementation(self, batch: ForwardBatch, inference_args: InferenceArgs) -> ForwardBatch:
|
||||
# "custom logic for HunYuan latent preparation"
|
||||
# pass
|
||||
|
||||
|
||||
# class hunyuanloader(PipelineLoader):
|
||||
# def load_components(self, inference_args: InferenceArgs):
|
||||
# pass
|
||||
|
||||
@dataclass
|
||||
class DiffusionPipelineOutput(BaseOutput):
|
||||
videos: Union[torch.Tensor, np.ndarray]
|
||||
|
||||
|
||||
class HunyuanVideoPipeline(ComposedPipelineBase):
|
||||
|
||||
def initialize_encoders(self, modules: Dict[str, Any], inference_args: InferenceArgs):
|
||||
self.initialize_encoders_v1(modules, inference_args)
|
||||
|
||||
|
||||
def initialize_encoders_v1(self, modules: Dict[str, Any], inference_args: InferenceArgs):
|
||||
"""
|
||||
Initialize the encoders. Will remove the encoders/tokenizers modules from the
|
||||
modules. Will add the TextEncoder or ImageEncoder to the modules.
|
||||
"""
|
||||
from fastvideo.v1.models.text_encoder import TextEncoder
|
||||
|
||||
crop_start = PROMPT_TEMPLATE["video"].get("crop_start", 0)
|
||||
|
||||
|
||||
max_length = inference_args.text_len + crop_start
|
||||
|
||||
# prompt_template
|
||||
prompt_template = PROMPT_TEMPLATE["image"]
|
||||
|
||||
# prompt_template_video
|
||||
prompt_template_video = PROMPT_TEMPLATE["video"]
|
||||
|
||||
encoder_1 = modules.pop("text_encoder")
|
||||
assert encoder_1 is not None, "Text encoder is not found"
|
||||
encoder_1.to(inference_args.device)
|
||||
encoder_1.to(dtype=PRECISION_TO_TYPE[inference_args.text_encoder_precision])
|
||||
encoder_1.requires_grad_(False)
|
||||
|
||||
tokenizer_1 = modules.pop("tokenizer")
|
||||
assert tokenizer_1 is not None, "Tokenizer is not found"
|
||||
|
||||
text_encoder = TextEncoder(
|
||||
text_encoder=encoder_1,
|
||||
tokenizer=tokenizer_1,
|
||||
# text_encoder_type="text_encoder"
|
||||
max_length=max_length,
|
||||
# text_encoder_precision=inference_args.text_encoder_precision,
|
||||
prompt_template=prompt_template,
|
||||
prompt_template_video=prompt_template_video,
|
||||
hidden_state_skip_layer=inference_args.hidden_state_skip_layer,
|
||||
apply_final_norm=False,
|
||||
device=inference_args.device if not inference_args.use_cpu_offload else "cpu",
|
||||
)
|
||||
|
||||
encoder_2 = modules.pop("text_encoder_2")
|
||||
assert encoder_2 is not None, "Text encoder 2 is not found"
|
||||
encoder_2.to(inference_args.device)
|
||||
encoder_2.to(dtype=PRECISION_TO_TYPE[inference_args.text_encoder_precision])
|
||||
encoder_2.requires_grad_(False)
|
||||
|
||||
tokenizer_2 = modules.pop("tokenizer_2")
|
||||
assert tokenizer_2 is not None, "Tokenizer 2 is not found"
|
||||
|
||||
text_encoder_2 = TextEncoder(
|
||||
text_encoder=encoder_2,
|
||||
tokenizer=tokenizer_2,
|
||||
# text_encoder_type="text_encoder_2",
|
||||
max_length=inference_args.text_len_2,
|
||||
# text_encoder_precision=inference_args.text_encoder_precision,
|
||||
device=inference_args.device if not inference_args.use_cpu_offload else "cpu",
|
||||
)
|
||||
modules["text_encoder"] = text_encoder
|
||||
modules["text_encoder_2"] = text_encoder_2
|
||||
|
||||
|
||||
def setup_pipeline(self, inference_args: InferenceArgs):
|
||||
self.add_stage("input_validation_stage",
|
||||
InputValidationStage())
|
||||
self.add_stage("prompt_encoding_stage_primary",
|
||||
PromptEncodingStage(is_secondary=False))
|
||||
self.add_stage("prompt_encoding_stage_secondary",
|
||||
PromptEncodingStage(is_secondary=True))
|
||||
self.add_stage("conditioning_stage",
|
||||
ConditioningStage())
|
||||
self.add_stage("timestep_preparation_stage",
|
||||
TimestepPreparationStage())
|
||||
self.add_stage("latent_preparation_stage",
|
||||
LatentPreparationStage())
|
||||
self.add_stage("denoising_stage",
|
||||
DenoisingStage())
|
||||
self.add_stage("decoding_stage",
|
||||
DecodingStage())
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
assert len(self._stages) > 0, "Pipeline stages are not set"
|
||||
assert len(self._modules) > 0, "Pipeline modules are not set"
|
||||
|
||||
|
||||
vae_scale_factor = 2**(len(self.vae.block_out_channels) - 1)
|
||||
inference_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
self.image_processor = VaeImageProcessor(vae_scale_factor=vae_scale_factor)
|
||||
self.register_modules({"image_processor": self.image_processor})
|
||||
|
||||
|
||||
num_channels_latents = self.transformer.in_channels
|
||||
inference_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
# TODO
|
||||
def adjust_video_length(self, batch: ForwardBatch, inference_args: InferenceArgs):
|
||||
"""Adjust video length based on VAE version"""
|
||||
video_length = batch.num_frames
|
||||
batch.num_frames = (video_length - 1) // 4 + 1
|
||||
return batch
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, inference_args: InferenceArgs) -> ForwardBatch:
|
||||
logger.info(f"Running pipeline stages: {self._stage_name_mapping.keys()}")
|
||||
logger.info(f"Batch: {batch}")
|
||||
# for stage in self._stages:
|
||||
# batch = stage(batch, inference_args)
|
||||
|
||||
# or
|
||||
|
||||
batch = self.input_validation_stage(batch, inference_args)
|
||||
batch = self.prompt_encoding_stage_primary(batch, inference_args)
|
||||
batch = self.prompt_encoding_stage_secondary(batch, inference_args)
|
||||
batch = self.conditioning_stage(batch, inference_args)
|
||||
batch = self.timestep_preparation_stage(batch, inference_args)
|
||||
|
||||
# custom logic
|
||||
batch = self.adjust_video_length(batch, inference_args)
|
||||
|
||||
batch = self.latent_preparation_stage(batch, inference_args)
|
||||
batch = self.denoising_stage(batch, inference_args)
|
||||
batch = self.decoding_stage(batch, inference_args)
|
||||
|
||||
return DiffusionPipelineOutput(videos=batch.videos)
|
||||
|
||||
EntryClass = HunyuanVideoPipeline
|
||||
@@ -0,0 +1,105 @@
|
||||
"""
|
||||
Data structures for functional pipeline processing.
|
||||
|
||||
This module defines the dataclasses used to pass state between pipeline components
|
||||
in a functional manner, reducing the need for explicit parameter passing.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, List, Optional, Union, Tuple, Callable
|
||||
import torch
|
||||
|
||||
|
||||
@dataclass
|
||||
class ForwardBatch:
|
||||
"""
|
||||
Complete state passed through the pipeline execution.
|
||||
|
||||
This dataclass contains all information needed during the diffusion pipeline
|
||||
execution, allowing methods to update specific components without needing
|
||||
to manage numerous individual parameters.
|
||||
"""
|
||||
# TODO(will): double check that args are separate from inference_args
|
||||
# properly. Also maybe think about providing an abstraction for pipeline
|
||||
# specific arguments.
|
||||
data_type: str
|
||||
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None
|
||||
mask_strategy: Optional[Dict[str, List[str]]] = None
|
||||
|
||||
# Text inputs
|
||||
prompt: Optional[Union[str, List[str]]] = None
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None
|
||||
|
||||
# Primary encoder embeddings
|
||||
prompt_embeds: Optional[torch.Tensor] = None
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None
|
||||
attention_mask: Optional[torch.Tensor] = None
|
||||
negative_attention_mask: Optional[torch.Tensor] = None
|
||||
|
||||
# Secondary encoder embeddings (for dual-encoder models)
|
||||
prompt_embeds_2: Optional[torch.Tensor] = None
|
||||
negative_prompt_embeds_2: Optional[torch.Tensor] = None
|
||||
attention_mask_2: Optional[torch.Tensor] = None
|
||||
negative_attention_mask_2: Optional[torch.Tensor] = None
|
||||
|
||||
# Additional text-related parameters
|
||||
max_sequence_length: Optional[int] = None
|
||||
prompt_template: Optional[Dict[str, Any]] = None
|
||||
do_classifier_free_guidance: bool = False
|
||||
|
||||
# Batch info
|
||||
batch_size: Optional[int] = None
|
||||
num_videos_per_prompt: int = 1
|
||||
|
||||
# Tracking if embeddings are already processed
|
||||
is_prompt_processed: bool = False
|
||||
|
||||
|
||||
# Latent tensors
|
||||
latents: Optional[torch.Tensor] = None
|
||||
noise_pred: Optional[torch.Tensor] = None
|
||||
|
||||
# Latent dimensions
|
||||
num_channels_latents: Optional[int] = None
|
||||
height_latents: Optional[int] = None
|
||||
width_latents: Optional[int] = None
|
||||
num_frames: int = 1 # Default for image models
|
||||
|
||||
# Original dimensions (before VAE scaling)
|
||||
height: Optional[int] = None
|
||||
width: Optional[int] = None
|
||||
|
||||
# Timesteps
|
||||
timesteps: Optional[torch.Tensor] = None
|
||||
timestep: Optional[Union[torch.Tensor, float, int]] = None
|
||||
step_index: Optional[int] = None
|
||||
|
||||
# Scheduler parameters
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
eta: float = 0.0
|
||||
sigmas: Optional[List[float]] = None
|
||||
|
||||
n_tokens: Optional[int] = None
|
||||
|
||||
# Other parameters that may be needed by specific schedulers
|
||||
extra_step_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# Component modules (populated by the pipeline)
|
||||
modules: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# Final output (after pipeline completion)
|
||||
output: Any = None
|
||||
|
||||
# Extra parameters that might be needed by specific pipeline implementations
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
device: torch.device = field(default_factory=lambda: torch.device("cuda"))
|
||||
|
||||
def __post_init__(self):
|
||||
"""Initialize dependent fields after dataclass initialization."""
|
||||
|
||||
# Set do_classifier_free_guidance based on guidance scale and negative prompt
|
||||
if self.guidance_scale > 1.0:
|
||||
self.do_classifier_free_guidance = True
|
||||
@@ -0,0 +1,91 @@
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.6.4.post1/vllm/model_executor/models/registry.py
|
||||
# and https://github.com/sgl-project/sglang/blob/v0.4.3/python/sglang/srt/models/registry.py
|
||||
|
||||
import importlib
|
||||
import pkgutil
|
||||
from dataclasses import dataclass, field
|
||||
from functools import lru_cache
|
||||
from typing import AbstractSet, Dict, List, Optional, Tuple, Type, Union
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PipelineRegistry:
|
||||
# Keyed by pipeline_arch
|
||||
pipelines: Dict[str, Union[Type[nn.Module], str]] = field(default_factory=dict)
|
||||
|
||||
def get_supported_archs(self) -> AbstractSet[str]:
|
||||
return self.pipelines.keys()
|
||||
|
||||
def _raise_for_unsupported(self, architectures: List[str]):
|
||||
all_supported_archs = self.get_supported_archs()
|
||||
|
||||
if any(arch in all_supported_archs for arch in architectures):
|
||||
raise ValueError(
|
||||
f"Pipeline architectures {architectures} failed "
|
||||
"to be inspected. Please check the logs for more details."
|
||||
)
|
||||
|
||||
raise ValueError(
|
||||
f"Pipeline architectures {architectures} are not supported for now. "
|
||||
f"Supported architectures: {all_supported_archs}"
|
||||
)
|
||||
|
||||
def _try_load_pipeline_cls(self, pipeline_arch: str) -> Optional[Type[nn.Module]]:
|
||||
if pipeline_arch not in self.pipelines:
|
||||
return None
|
||||
|
||||
return self.pipelines[pipeline_arch]
|
||||
|
||||
def resolve_pipeline_cls(
|
||||
self,
|
||||
architecture: str,
|
||||
) -> Tuple[Type[nn.Module], str]:
|
||||
if not architecture:
|
||||
logger.warning("No pipeline architecture is specified")
|
||||
|
||||
pipeline_cls = self._try_load_pipeline_cls(architecture)
|
||||
if pipeline_cls is not None:
|
||||
return (pipeline_cls, architecture)
|
||||
|
||||
return self._raise_for_unsupported(architecture)
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def import_pipeline_classes():
|
||||
pipeline_arch_name_to_cls = {}
|
||||
package_name = "fastvideo.v1.pipelines.implementations"
|
||||
package = importlib.import_module(package_name)
|
||||
for _, name, ispkg in pkgutil.iter_modules(package.__path__, package_name + "."):
|
||||
if ispkg:
|
||||
try:
|
||||
module = importlib.import_module(name)
|
||||
except Exception as e:
|
||||
logger.warning(f"Ignore import error when loading {name}. " f"{e}")
|
||||
continue
|
||||
if hasattr(module, "EntryClass"):
|
||||
entry = module.EntryClass
|
||||
print(entry)
|
||||
print(entry.__name__)
|
||||
if isinstance(
|
||||
entry, list
|
||||
): # To support multiple pipeline classes in one module
|
||||
for tmp in entry:
|
||||
assert (
|
||||
tmp.__name__ not in pipeline_arch_name_to_cls
|
||||
), f"Duplicated pipeline implementation for {tmp.__name__}"
|
||||
pipeline_arch_name_to_cls[tmp.__name__] = tmp
|
||||
else:
|
||||
assert (
|
||||
entry.__name__ not in pipeline_arch_name_to_cls
|
||||
), f"Duplicated pipeline implementation for {entry.__name__}"
|
||||
pipeline_arch_name_to_cls[entry.__name__] = entry
|
||||
return pipeline_arch_name_to_cls
|
||||
|
||||
|
||||
PipelineRegistry = _PipelineRegistry(import_pipeline_classes())
|
||||
@@ -0,0 +1,28 @@
|
||||
"""
|
||||
Pipeline stages for diffusion models.
|
||||
|
||||
This package contains the various stages that can be composed to create
|
||||
complete diffusion pipelines.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.input_validation import InputValidationStage
|
||||
from fastvideo.v1.pipelines.stages.prompt_encoding import PromptEncodingStage
|
||||
from fastvideo.v1.pipelines.stages.timestep_preparation import TimestepPreparationStage
|
||||
from fastvideo.v1.pipelines.stages.latent_preparation import LatentPreparationStage
|
||||
from fastvideo.v1.pipelines.stages.conditioning import ConditioningStage
|
||||
from fastvideo.v1.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.v1.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.v1.pipelines.stages.post_processing import PostProcessingStage
|
||||
|
||||
__all__ = [
|
||||
"PipelineStage",
|
||||
"InputValidationStage",
|
||||
"PromptEncodingStage",
|
||||
"TimestepPreparationStage",
|
||||
"LatentPreparationStage",
|
||||
"ConditioningStage",
|
||||
"DenoisingStage",
|
||||
"DecodingStage",
|
||||
"PostProcessingStage",
|
||||
]
|
||||
@@ -0,0 +1,125 @@
|
||||
"""
|
||||
Base classes for pipeline stages.
|
||||
|
||||
This module defines the abstract base classes for pipeline stages that can be
|
||||
composed to create complete diffusion pipelines.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
import torch
|
||||
import time
|
||||
import traceback
|
||||
from typing import Dict, Any
|
||||
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PipelineStage(ABC):
|
||||
"""
|
||||
Abstract base class for all pipeline stages.
|
||||
|
||||
A pipeline stage represents a discrete step in the diffusion process that can be
|
||||
composed with other stages to create a complete pipeline. Each stage is responsible
|
||||
for a specific part of the process, such as prompt encoding, latent preparation, etc.
|
||||
"""
|
||||
|
||||
def __init__(self, enable_logging: bool = False):
|
||||
"""
|
||||
Initialize the pipeline stage.
|
||||
|
||||
Args:
|
||||
enable_logging: Whether to enable logging for this stage.
|
||||
"""
|
||||
self._enable_logging = enable_logging
|
||||
self._stage_name = self.__class__.__name__
|
||||
self._logger = init_logger(f"fastvideo.v1.pipelines.stages.{self._stage_name}")
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
"""Get the device for this stage."""
|
||||
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
def set_logging(self, enable: bool):
|
||||
"""
|
||||
Enable or disable logging for this stage.
|
||||
|
||||
Args:
|
||||
enable: Whether to enable logging.
|
||||
"""
|
||||
self._enable_logging = enable
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Execute the stage's processing on the batch with optional logging.
|
||||
Should not be overridden by subclasses.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The updated batch information after this stage's processing.
|
||||
"""
|
||||
if self._enable_logging:
|
||||
self._logger.info(f"[{self._stage_name}] Starting execution")
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
# Call the actual implementation
|
||||
result = self._call_implementation(batch, inference_args)
|
||||
|
||||
execution_time = time.time() - start_time
|
||||
self._logger.info(f"[{self._stage_name}] Execution completed in {execution_time * 1000:.2f} ms")
|
||||
|
||||
return result
|
||||
except Exception as e:
|
||||
execution_time = time.time() - start_time
|
||||
self._logger.error(f"[{self._stage_name}] Error during execution after {execution_time * 1000:.2f} ms: {e}")
|
||||
self._logger.error(f"[{self._stage_name}] Traceback: {traceback.format_exc()}")
|
||||
|
||||
# Re-raise the exception
|
||||
raise
|
||||
else:
|
||||
# Just call the implementation directly if logging is disabled
|
||||
return self._call_implementation(batch, inference_args)
|
||||
|
||||
@abstractmethod
|
||||
def _call_implementation(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Actual implementation of the stage's processing.
|
||||
|
||||
This method should be implemented by subclasses to provide the actual
|
||||
processing logic for the stage.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The updated batch information after this stage's processing.
|
||||
"""
|
||||
pass
|
||||
|
||||
def register_modules(self, modules: Dict[str, Any]):
|
||||
"""
|
||||
Register modules needed by this stage.
|
||||
|
||||
Args:
|
||||
modules: The modules to register.
|
||||
"""
|
||||
for name, module in modules.items():
|
||||
if self._enable_logging:
|
||||
self._logger.debug(f"[{self._stage_name}] Registering module: {name}")
|
||||
setattr(self, name, module)
|
||||
@@ -0,0 +1,70 @@
|
||||
"""
|
||||
Conditioning stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from typing import Optional
|
||||
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ConditioningStage(PipelineStage):
|
||||
"""
|
||||
Stage for applying conditioning to the diffusion process.
|
||||
|
||||
This stage handles the application of conditioning, such as classifier-free guidance,
|
||||
to the diffusion process.
|
||||
"""
|
||||
|
||||
def _call_implementation(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Apply conditioning to the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with applied conditioning.
|
||||
"""
|
||||
if not batch.do_classifier_free_guidance:
|
||||
return batch
|
||||
|
||||
logger.info(f"batch.negative_prompt_embeds: {batch.negative_prompt_embeds}")
|
||||
logger.info(f"do_classifier_free_guidance: {batch.do_classifier_free_guidance}")
|
||||
logger.info(f"cfg_scale: {batch.guidance_scale}")
|
||||
|
||||
# Ensure negative prompt embeddings are available
|
||||
assert batch.negative_prompt_embeds is not None, (
|
||||
"Negative prompt embeddings are required for classifier-free guidance"
|
||||
)
|
||||
|
||||
# Concatenate primary embeddings and masks
|
||||
batch.prompt_embeds = torch.cat(
|
||||
[batch.negative_prompt_embeds, batch.prompt_embeds]
|
||||
)
|
||||
if batch.attention_mask is not None:
|
||||
batch.attention_mask = torch.cat(
|
||||
[batch.negative_attention_mask, batch.attention_mask]
|
||||
)
|
||||
|
||||
# Concatenate secondary embeddings and masks if present
|
||||
if batch.prompt_embeds_2 is not None:
|
||||
batch.prompt_embeds_2 = torch.cat(
|
||||
[batch.negative_prompt_embeds_2, batch.prompt_embeds_2]
|
||||
)
|
||||
if batch.attention_mask_2 is not None:
|
||||
batch.attention_mask_2 = torch.cat(
|
||||
[batch.negative_attention_mask_2, batch.attention_mask_2]
|
||||
)
|
||||
|
||||
return batch
|
||||
@@ -0,0 +1,77 @@
|
||||
"""
|
||||
Decoding stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
from fastvideo.v1.logger import init_logger
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class DecodingStage(PipelineStage):
|
||||
"""
|
||||
Stage for decoding latent representations into pixel space.
|
||||
|
||||
This stage handles the decoding of latent representations into the final
|
||||
output format (e.g., pixel values).
|
||||
"""
|
||||
|
||||
def _call_implementation(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Decode latent representations into pixel space.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with decoded outputs.
|
||||
"""
|
||||
latents = batch.latents
|
||||
|
||||
# Skip decoding if output type is latent
|
||||
if inference_args.output_type == "latent":
|
||||
image = latents
|
||||
else:
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32) and not inference_args.disable_autocast
|
||||
|
||||
|
||||
# Apply scaling/shifting if needed
|
||||
if (hasattr(self.vae.config, "shift_factor") and self.vae.config.shift_factor):
|
||||
latents = (latents / self.vae.config.scaling_factor + self.vae.config.shift_factor)
|
||||
else:
|
||||
latents = latents / self.vae.config.scaling_factor
|
||||
|
||||
# Decode latents
|
||||
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
|
||||
if inference_args.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if inference_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
|
||||
# Convert to CPU float32 for compatibility
|
||||
image = image.cpu().float()
|
||||
|
||||
# Update batch with decoded image
|
||||
batch.videos = image
|
||||
|
||||
# Offload models if needed
|
||||
if hasattr(self, 'maybe_free_model_hooks'):
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
return batch
|
||||
@@ -0,0 +1,210 @@
|
||||
"""
|
||||
Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
import torch
|
||||
from typing import Optional, Dict, Any, List
|
||||
from tqdm.auto import tqdm
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
# TODO(will-refactor): change this to fastvideo.distributed
|
||||
from fastvideo.v1.distributed import get_sequence_model_parallel_world_size, get_sequence_model_parallel_rank
|
||||
from fastvideo.v1.distributed.communication_op import sequence_model_parallel_all_gather
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class DenoisingStage(PipelineStage):
|
||||
"""
|
||||
Stage for running the denoising loop in diffusion pipelines.
|
||||
|
||||
This stage handles the iterative denoising process that transforms
|
||||
the initial noise into the final output.
|
||||
"""
|
||||
|
||||
def _call_implementation(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Run the denoising loop.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
# Prepare extra step kwargs for scheduler
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.step,
|
||||
{
|
||||
"generator": batch.generator,
|
||||
"eta": batch.eta
|
||||
},
|
||||
)
|
||||
|
||||
# Setup precision and autocast settings
|
||||
target_dtype = PRECISION_TO_TYPE[inference_args.precision]
|
||||
autocast_enabled = (target_dtype != torch.float32) and not inference_args.disable_autocast
|
||||
|
||||
# Handle sequence parallelism if enabled
|
||||
world_size, rank = get_sequence_model_parallel_world_size(), get_sequence_model_parallel_rank()
|
||||
sp_group = True if world_size > 1 else False
|
||||
if sp_group:
|
||||
latents = rearrange(batch.latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
|
||||
latents = latents[:, :, rank, :, :, :]
|
||||
batch.latents = latents
|
||||
|
||||
# Get timesteps and calculate warmup steps
|
||||
timesteps = batch.timesteps
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
|
||||
# Create 3D list for mask strategy
|
||||
def dict_to_3d_list(mask_strategy, t_max=50, l_max=60, h_max=24):
|
||||
result = [[[None for _ in range(h_max)] for _ in range(l_max)] for _ in range(t_max)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
for key, value in mask_strategy.items():
|
||||
t, l, h = map(int, key.split('_'))
|
||||
result[t][l][h] = value
|
||||
return result
|
||||
|
||||
mask_strategy = dict_to_3d_list(batch.mask_strategy)
|
||||
|
||||
# Get latents and embeddings
|
||||
latents = batch.latents
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
prompt_embeds_2 = batch.prompt_embeds_2
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
# Skip if interrupted
|
||||
if hasattr(self, 'interrupt') and self.interrupt:
|
||||
continue
|
||||
|
||||
# Expand latents for classifier-free guidance
|
||||
latent_model_input = (torch.cat([latents] * 2) if batch.do_classifier_free_guidance else latents)
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
|
||||
# Prepare inputs for transformer
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (torch.tensor(
|
||||
[inference_args.embedded_cfg_scale] * latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=batch.device,
|
||||
).to(target_dtype) * 1000.0 if inference_args.embedded_cfg_scale is not None else None)
|
||||
|
||||
# Predict noise residual
|
||||
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
|
||||
# Prepare encoder hidden states
|
||||
if prompt_embeds_2 is not None and prompt_embeds_2.shape[-1] != prompt_embeds.shape[-1]:
|
||||
prompt_embeds_2 = torch.nn.functional.pad(
|
||||
prompt_embeds_2,
|
||||
(0, prompt_embeds.shape[2] - prompt_embeds_2.shape[1]),
|
||||
value=0,
|
||||
).unsqueeze(1)
|
||||
encoder_hidden_states = torch.cat([prompt_embeds_2, prompt_embeds], dim=1) if prompt_embeds_2 is not None else prompt_embeds
|
||||
|
||||
# Run transformer
|
||||
noise_pred = self.transformer(
|
||||
latent_model_input,
|
||||
encoder_hidden_states,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
)
|
||||
|
||||
# Apply guidance
|
||||
if batch.do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + batch.guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
# Apply guidance rescale if needed
|
||||
if batch.guidance_rescale > 0.0:
|
||||
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
|
||||
noise_pred = self.rescale_noise_cfg(
|
||||
noise_pred,
|
||||
noise_pred_text,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
)
|
||||
|
||||
# Compute the previous noisy sample
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
if progress_bar is not None:
|
||||
progress_bar.update()
|
||||
|
||||
# Gather results if using sequence parallelism
|
||||
if sp_group:
|
||||
latents = sequence_model_parallel_all_gather(latents, dim=2)
|
||||
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
return batch
|
||||
|
||||
def prepare_extra_func_kwargs(self, func, kwargs):
|
||||
"""
|
||||
Prepare extra kwargs for the scheduler step.
|
||||
|
||||
Args:
|
||||
func: The function to prepare kwargs for.
|
||||
kwargs: The kwargs to prepare.
|
||||
|
||||
Returns:
|
||||
The prepared kwargs.
|
||||
"""
|
||||
extra_step_kwargs = {}
|
||||
for k, v in kwargs.items():
|
||||
accepts = k in set(inspect.signature(func).parameters.keys())
|
||||
if accepts:
|
||||
extra_step_kwargs[k] = v
|
||||
return extra_step_kwargs
|
||||
|
||||
def progress_bar(self, iterable=None, total=None):
|
||||
"""
|
||||
Create a progress bar for the denoising process.
|
||||
|
||||
Args:
|
||||
iterable: The iterable to iterate over.
|
||||
total: The total number of items.
|
||||
|
||||
Returns:
|
||||
A tqdm progress bar.
|
||||
"""
|
||||
return tqdm(iterable=iterable, total=total)
|
||||
|
||||
def rescale_noise_cfg(self, noise_cfg, noise_pred_text, guidance_rescale=0.0):
|
||||
"""
|
||||
Rescale noise prediction according to guidance_rescale.
|
||||
|
||||
Based on findings of "Common Diffusion Noise Schedules and Sample Steps are Flawed"
|
||||
(https://arxiv.org/pdf/2305.08891.pdf), Section 3.4.
|
||||
|
||||
Args:
|
||||
noise_cfg: The noise prediction with guidance.
|
||||
noise_pred_text: The text-conditioned noise prediction.
|
||||
guidance_rescale: The guidance rescale factor.
|
||||
|
||||
Returns:
|
||||
The rescaled noise prediction.
|
||||
"""
|
||||
std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True)
|
||||
std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)
|
||||
# Rescale the results from guidance (fixes overexposure)
|
||||
noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
|
||||
# Mix with the original results from guidance by factor guidance_rescale
|
||||
noise_cfg = (guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg)
|
||||
return noise_cfg
|
||||
@@ -0,0 +1,91 @@
|
||||
"""
|
||||
Input validation stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
from typing import Optional, Union, List
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
import random
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class InputValidationStage(PipelineStage):
|
||||
"""
|
||||
Stage for validating and preparing inputs for diffusion pipelines.
|
||||
|
||||
This stage validates that all required inputs are present and properly formatted
|
||||
before proceeding with the diffusion process.
|
||||
"""
|
||||
|
||||
def _generate_seeds(self, batch: ForwardBatch, inference_args: InferenceArgs):
|
||||
"""Generate seeds for the inference"""
|
||||
seed = inference_args.seed
|
||||
num_videos_per_prompt = inference_args.num_videos
|
||||
|
||||
|
||||
seeds = [seed + i for i in range(num_videos_per_prompt)]
|
||||
batch.seeds = seeds
|
||||
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
|
||||
batch.generator = [torch.Generator("cpu").manual_seed(seed) for seed in seeds]
|
||||
|
||||
def _call_implementation(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Validate and prepare inputs.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The validated batch information.
|
||||
"""
|
||||
self._generate_seeds(batch, inference_args)
|
||||
|
||||
# Ensure prompt is properly formatted
|
||||
if batch.prompt is None and batch.prompt_embeds is None:
|
||||
raise ValueError("Either `prompt` or `prompt_embeds` must be provided")
|
||||
|
||||
# Ensure negative prompt is properly formatted if using classifier-free guidance
|
||||
if batch.do_classifier_free_guidance:
|
||||
if batch.negative_prompt is None and batch.negative_prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"For classifier-free guidance, either `negative_prompt` or "
|
||||
"`negative_prompt_embeds` must be provided"
|
||||
)
|
||||
|
||||
# Validate height and width
|
||||
if batch.height % 8 != 0 or batch.width % 8 != 0:
|
||||
raise ValueError(
|
||||
f"Height and width must be divisible by 8 but are {batch.height} and {batch.width}."
|
||||
)
|
||||
|
||||
# Validate number of inference steps
|
||||
if batch.num_inference_steps <= 0:
|
||||
raise ValueError(
|
||||
f"Number of inference steps must be positive, but got {batch.num_inference_steps}"
|
||||
)
|
||||
|
||||
# Validate guidance scale if using classifier-free guidance
|
||||
if batch.do_classifier_free_guidance and batch.guidance_scale <= 0:
|
||||
raise ValueError(
|
||||
f"Guidance scale must be positive, but got {batch.guidance_scale}"
|
||||
)
|
||||
|
||||
# Set device if not already set
|
||||
if batch.device is None:
|
||||
batch.device = self.device
|
||||
|
||||
# Set data type if not already set
|
||||
if batch.data_type is None:
|
||||
batch.data_type = inference_args.precision
|
||||
|
||||
return batch
|
||||
@@ -0,0 +1,113 @@
|
||||
"""
|
||||
Latent preparation stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LatentPreparationStage(PipelineStage):
|
||||
"""
|
||||
Stage for preparing initial latent variables for the diffusion process.
|
||||
|
||||
This stage handles the preparation of the initial latent variables that will be
|
||||
denoised during the diffusion process.
|
||||
"""
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
# TODO(will): this is a hack to get the vae scale factor. Check if this
|
||||
# is only needed for hunyuan
|
||||
|
||||
def _call_implementation(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Prepare initial latent variables for the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with prepared latent variables.
|
||||
"""
|
||||
# Determine batch size
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
elif batch.prompt is not None:
|
||||
batch_size = 1
|
||||
else:
|
||||
batch_size = batch.prompt_embeds.shape[0]
|
||||
|
||||
# Adjust batch size for number of videos per prompt
|
||||
batch_size *= batch.num_videos_per_prompt
|
||||
|
||||
# Get required parameters
|
||||
dtype = batch.prompt_embeds.dtype
|
||||
device = batch.device
|
||||
generator = batch.generator
|
||||
latents = batch.latents
|
||||
num_frames = batch.num_frames
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
|
||||
# Calculate latent shape
|
||||
shape = (
|
||||
batch_size,
|
||||
inference_args.num_channels_latents,
|
||||
num_frames,
|
||||
int(height) // inference_args.vae_scale_factor,
|
||||
int(width) // inference_args.vae_scale_factor,
|
||||
)
|
||||
|
||||
# Validate generator if it's a list
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
# Generate or use provided latents
|
||||
if latents is None:
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
else:
|
||||
latents = latents.to(device)
|
||||
|
||||
# Scale the initial noise if needed
|
||||
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
|
||||
# Update batch with prepared latents
|
||||
batch.latents = latents
|
||||
|
||||
# Adjust video length based on VAE version if needed
|
||||
if hasattr(self, 'adjust_video_length'):
|
||||
batch = self.adjust_video_length(batch, inference_args)
|
||||
|
||||
return batch
|
||||
|
||||
def adjust_video_length(self, batch: ForwardBatch, inference_args: InferenceArgs) -> ForwardBatch:
|
||||
"""
|
||||
Adjust video length based on VAE version.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with adjusted video length.
|
||||
"""
|
||||
video_length = batch.num_frames
|
||||
# TODO
|
||||
batch.num_frames = (video_length - 1) // 4 + 1
|
||||
return batch
|
||||
@@ -0,0 +1,61 @@
|
||||
"""
|
||||
Post-processing stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
from typing import Optional, Union, List, Tuple
|
||||
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from diffusers.utils import BaseOutput
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PostProcessingStage(PipelineStage):
|
||||
"""
|
||||
Stage for post-processing the decoded outputs.
|
||||
|
||||
This stage handles any final processing needed on the decoded outputs,
|
||||
such as format conversion, normalization, etc.
|
||||
"""
|
||||
|
||||
def _call_implementation(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Apply post-processing to the results.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with post-processed outputs.
|
||||
"""
|
||||
videos = batch.videos
|
||||
|
||||
# Convert to numpy if requested
|
||||
if inference_args.output_type == "numpy":
|
||||
videos = videos.numpy()
|
||||
|
||||
# Create output object
|
||||
output = DiffusionPipelineOutput(videos=videos)
|
||||
batch.output = output
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
class DiffusionPipelineOutput(BaseOutput):
|
||||
"""
|
||||
Output class for diffusion pipelines.
|
||||
|
||||
Args:
|
||||
videos: The generated videos.
|
||||
"""
|
||||
videos: Union[torch.Tensor, np.ndarray]
|
||||
@@ -0,0 +1,113 @@
|
||||
"""
|
||||
Prompt encoding stages for diffusion pipelines.
|
||||
|
||||
This module contains implementations of prompt encoding stages for diffusion pipelines.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Union, Tuple
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.pipelines.stages import PipelineStage
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PromptEncodingStage(PipelineStage):
|
||||
"""
|
||||
Stage for encoding text prompts into embeddings for diffusion models.
|
||||
|
||||
This stage handles the encoding of text prompts into the embedding space
|
||||
expected by the diffusion model.
|
||||
"""
|
||||
|
||||
def __init__(self, enable_logging: bool = False, is_secondary: bool = False):
|
||||
"""
|
||||
Initialize the prompt encoding stage.
|
||||
|
||||
Args:
|
||||
enable_logging: Whether to enable logging for this stage.
|
||||
is_secondary: Whether this is a secondary text encoder.
|
||||
"""
|
||||
super().__init__(enable_logging=enable_logging)
|
||||
self.is_secondary = is_secondary
|
||||
|
||||
def _call_implementation(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if self.is_secondary:
|
||||
assert self.text_encoder_2 is not None, "Secondary text encoder is not set"
|
||||
text_encoder = self.text_encoder_2
|
||||
else:
|
||||
text_encoder = self.text_encoder
|
||||
|
||||
prompt: Union[str, List[str]] = batch.prompt
|
||||
device: torch.device = batch.device
|
||||
num_videos_per_prompt: int = batch.num_videos_per_prompt
|
||||
data_type: str = batch.data_type
|
||||
|
||||
# Get the right prompt embeds and attention masks based on whether this is primary or secondary
|
||||
if self.is_secondary:
|
||||
prompt_embeds = batch.prompt_embeds_2
|
||||
attention_mask = batch.attention_mask_2
|
||||
else:
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
attention_mask = batch.attention_mask
|
||||
|
||||
if prompt_embeds is None:
|
||||
# textual inversion: process multi-vector tokens if necessary
|
||||
# if isinstance(self, TextualInversionLoaderMixin):
|
||||
# prompt = self.maybe_convert_prompt(prompt, text_encoder.tokenizer)
|
||||
|
||||
text_inputs = text_encoder.text2tokens(prompt)
|
||||
prompt_outputs = text_encoder.encode(text_inputs, device=device)
|
||||
prompt_embeds = prompt_outputs.hidden_state
|
||||
# TODO(will): support clip_skip
|
||||
|
||||
if text_encoder is not None:
|
||||
# TODO(will-refactor): use text_encoder.dtype
|
||||
prompt_embeds_dtype = torch.float16
|
||||
elif self.transformer is not None:
|
||||
prompt_embeds_dtype = self.transformer.dtype
|
||||
else:
|
||||
prompt_embeds_dtype = prompt_embeds.dtype
|
||||
|
||||
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device)
|
||||
# print("prompt_embeds", type(prompt_embeds))
|
||||
# logger.info(f"prompt_embeds shape: {prompt_embeds.shape}")
|
||||
if prompt_embeds.ndim == 1:
|
||||
prompt_embeds = prompt_embeds.unsqueeze(0)
|
||||
|
||||
if prompt_embeds.ndim == 2:
|
||||
bs_embed, _ = prompt_embeds.shape
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt)
|
||||
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, -1)
|
||||
else:
|
||||
bs_embed, seq_len, _ = prompt_embeds.shape
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
# Set the appropriate attributes based on whether this is primary or secondary
|
||||
if self.is_secondary:
|
||||
batch.prompt_embeds_2 = prompt_embeds
|
||||
else:
|
||||
batch.prompt_embeds = prompt_embeds
|
||||
|
||||
|
||||
return batch
|
||||
@@ -0,0 +1,86 @@
|
||||
"""
|
||||
Timestep preparation stages for diffusion pipelines.
|
||||
|
||||
This module contains implementations of timestep preparation stages for diffusion pipelines.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Union, Tuple
|
||||
import torch
|
||||
import inspect
|
||||
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class TimestepPreparationStage(PipelineStage):
|
||||
"""
|
||||
Stage for preparing timesteps for the diffusion process.
|
||||
|
||||
This stage handles the preparation of the timestep sequence that will be used
|
||||
during the diffusion process.
|
||||
"""
|
||||
|
||||
def _call_implementation(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Prepare timesteps for the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with prepared timesteps.
|
||||
"""
|
||||
scheduler = self.scheduler
|
||||
device = batch.device
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
timesteps = batch.timesteps
|
||||
sigmas = batch.sigmas
|
||||
n_tokens = batch.n_tokens
|
||||
|
||||
# Prepare extra kwargs for set_timesteps
|
||||
extra_set_timesteps_kwargs = {}
|
||||
if n_tokens is not None and "n_tokens" in inspect.signature(scheduler.set_timesteps).parameters:
|
||||
extra_set_timesteps_kwargs["n_tokens"] = n_tokens
|
||||
|
||||
# Handle custom timesteps or sigmas
|
||||
if timesteps is not None and sigmas is not None:
|
||||
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
|
||||
|
||||
if timesteps is not None:
|
||||
accepts_timesteps = "timesteps" in inspect.signature(scheduler.set_timesteps).parameters
|
||||
if not accepts_timesteps:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" timestep schedules. Please check whether you are using the correct scheduler."
|
||||
)
|
||||
scheduler.set_timesteps(timesteps=timesteps, device=device, **extra_set_timesteps_kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
elif sigmas is not None:
|
||||
accept_sigmas = "sigmas" in inspect.signature(scheduler.set_timesteps).parameters
|
||||
if not accept_sigmas:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" sigmas schedules. Please check whether you are using the correct scheduler."
|
||||
)
|
||||
scheduler.set_timesteps(sigmas=sigmas, device=device, **extra_set_timesteps_kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
else:
|
||||
scheduler.set_timesteps(num_inference_steps, device=device, **extra_set_timesteps_kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
|
||||
# Update batch with prepared timesteps
|
||||
batch.timesteps = timesteps
|
||||
batch.num_inference_steps = num_inference_steps
|
||||
|
||||
return batch
|
||||
@@ -0,0 +1,102 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import logging
|
||||
import traceback
|
||||
from contextlib import suppress
|
||||
from itertools import chain
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from vllm.plugins import load_plugins_by_group
|
||||
from vllm.utils import resolve_obj_by_qualname
|
||||
|
||||
from .interface import _Backend # noqa: F401
|
||||
from .interface import Platform, PlatformEnum
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def cuda_platform_plugin() -> Optional[str]:
|
||||
is_cuda = False
|
||||
|
||||
try:
|
||||
from vllm.utils import import_pynvml
|
||||
pynvml = import_pynvml()
|
||||
pynvml.nvmlInit()
|
||||
try:
|
||||
# NOTE: Edge case: vllm cpu build on a GPU machine.
|
||||
# Third-party pynvml can be imported in cpu build,
|
||||
# we need to check if vllm is built with cpu too.
|
||||
# Otherwise, vllm will always activate cuda plugin
|
||||
# on a GPU machine, even if in a cpu build.
|
||||
is_cuda = (pynvml.nvmlDeviceGetCount() > 0)
|
||||
finally:
|
||||
pynvml.nvmlShutdown()
|
||||
except Exception as e:
|
||||
if "nvml" not in e.__class__.__name__.lower():
|
||||
# If the error is not related to NVML, re-raise it.
|
||||
raise e
|
||||
|
||||
# CUDA is supported on Jetson, but NVML may not be.
|
||||
import os
|
||||
|
||||
def cuda_is_jetson() -> bool:
|
||||
return os.path.isfile("/etc/nv_tegra_release") \
|
||||
or os.path.exists("/sys/class/tegra-firmware")
|
||||
|
||||
if cuda_is_jetson():
|
||||
is_cuda = True
|
||||
|
||||
return "vllm.platforms.cuda.CudaPlatform" if is_cuda else None
|
||||
|
||||
builtin_platform_plugins = {
|
||||
'cuda': cuda_platform_plugin,
|
||||
}
|
||||
|
||||
|
||||
def resolve_current_platform_cls_qualname() -> str:
|
||||
# TODO(will): if we need to support other platforms, we should consider if
|
||||
# vLLM's plugin architecture is suitable for our needs.
|
||||
platform_cls_qualname = builtin_platform_plugins['cuda']()
|
||||
return platform_cls_qualname
|
||||
|
||||
|
||||
|
||||
_current_platform = None
|
||||
_init_trace: str = ''
|
||||
|
||||
if TYPE_CHECKING:
|
||||
current_platform: Platform
|
||||
|
||||
|
||||
def __getattr__(name: str):
|
||||
if name == 'current_platform':
|
||||
# lazy init current_platform.
|
||||
# 1. out-of-tree platform plugins need `from vllm.platforms import
|
||||
# Platform` so that they can inherit `Platform` class. Therefore,
|
||||
# we cannot resolve `current_platform` during the import of
|
||||
# `vllm.platforms`.
|
||||
# 2. when users use out-of-tree platform plugins, they might run
|
||||
# `import vllm`, some vllm internal code might access
|
||||
# `current_platform` during the import, and we need to make sure
|
||||
# `current_platform` is only resolved after the plugins are loaded
|
||||
# (we have tests for this, if any developer violate this, they will
|
||||
# see the test failures).
|
||||
global _current_platform
|
||||
if _current_platform is None:
|
||||
platform_cls_qualname = resolve_current_platform_cls_qualname()
|
||||
_current_platform = resolve_obj_by_qualname(
|
||||
platform_cls_qualname)()
|
||||
global _init_trace
|
||||
_init_trace = "".join(traceback.format_stack())
|
||||
return _current_platform
|
||||
elif name in globals():
|
||||
return globals()[name]
|
||||
else:
|
||||
raise AttributeError(
|
||||
f"No attribute named '{name}' exists in {__name__}.")
|
||||
|
||||
|
||||
__all__ = [
|
||||
'Platform', 'PlatformEnum', 'current_platform',
|
||||
"_init_trace"
|
||||
]
|
||||
@@ -0,0 +1,273 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Code inside this file can safely assume cuda platform, e.g. importing
|
||||
pynvml. However, it should not initialize cuda context.
|
||||
"""
|
||||
|
||||
import os
|
||||
from functools import lru_cache, wraps
|
||||
from typing import (TYPE_CHECKING, Callable, List, Optional, Tuple, TypeVar,
|
||||
Union)
|
||||
|
||||
import torch
|
||||
from typing_extensions import ParamSpec
|
||||
|
||||
# import custom ops, trigger op registration
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import import_pynvml
|
||||
|
||||
from .interface import DeviceCapability, Platform, PlatformEnum, _Backend
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_P = ParamSpec("_P")
|
||||
_R = TypeVar("_R")
|
||||
|
||||
pynvml = import_pynvml()
|
||||
|
||||
# pytorch 2.5 uses cudnn sdpa by default, which will cause crash on some models
|
||||
# see https://github.com/huggingface/diffusers/issues/9704 for details
|
||||
torch.backends.cuda.enable_cudnn_sdp(False)
|
||||
|
||||
|
||||
def device_id_to_physical_device_id(device_id: int) -> int:
|
||||
if "CUDA_VISIBLE_DEVICES" in os.environ:
|
||||
device_ids = os.environ["CUDA_VISIBLE_DEVICES"].split(",")
|
||||
if device_ids == [""]:
|
||||
msg = (
|
||||
"CUDA_VISIBLE_DEVICES is set to empty string, which means"
|
||||
" GPU support is disabled. If you are using ray, please unset"
|
||||
" the environment variable `CUDA_VISIBLE_DEVICES` inside the"
|
||||
" worker/actor. "
|
||||
"Check https://github.com/vllm-project/vllm/issues/8402 for"
|
||||
" more information.")
|
||||
raise RuntimeError(msg)
|
||||
physical_device_id = device_ids[device_id]
|
||||
return int(physical_device_id)
|
||||
else:
|
||||
return device_id
|
||||
|
||||
|
||||
def with_nvml_context(fn: Callable[_P, _R]) -> Callable[_P, _R]:
|
||||
|
||||
@wraps(fn)
|
||||
def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _R:
|
||||
pynvml.nvmlInit()
|
||||
try:
|
||||
return fn(*args, **kwargs)
|
||||
finally:
|
||||
pynvml.nvmlShutdown()
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
class CudaPlatformBase(Platform):
|
||||
_enum = PlatformEnum.CUDA
|
||||
device_name: str = "cuda"
|
||||
device_type: str = "cuda"
|
||||
dispatch_key: str = "CUDA"
|
||||
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls,
|
||||
device_id: int = 0
|
||||
) -> Optional[DeviceCapability]:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: Optional[bool]) -> bool:
|
||||
if enforce_eager:
|
||||
logger.warning(
|
||||
"To see benefits of async output processing, enable CUDA "
|
||||
"graph. Since, enforce-eager is enabled, async output "
|
||||
"processor cannot be used")
|
||||
return False
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def is_full_nvlink(cls, device_ids: List[int]) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def log_warnings(cls):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def get_current_memory_usage(cls,
|
||||
device: Optional[torch.types.Device] = None
|
||||
) -> float:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
return torch.cuda.max_memory_allocated(device)
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend, head_size, dtype,
|
||||
kv_cache_dtype, block_size, use_v1,
|
||||
use_mla) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
return "fastvideo.v1.distributed.device_communicators.cuda_communicator.CudaCommunicator" # noqa
|
||||
|
||||
|
||||
# NVML utils
|
||||
# Note that NVML is not affected by `CUDA_VISIBLE_DEVICES`,
|
||||
# all the related functions work on real physical device ids.
|
||||
# the major benefit of using NVML is that it will not initialize CUDA
|
||||
class NvmlCudaPlatform(CudaPlatformBase):
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=8)
|
||||
@with_nvml_context
|
||||
def get_device_capability(cls,
|
||||
device_id: int = 0
|
||||
) -> Optional[DeviceCapability]:
|
||||
try:
|
||||
physical_device_id = device_id_to_physical_device_id(device_id)
|
||||
handle = pynvml.nvmlDeviceGetHandleByIndex(physical_device_id)
|
||||
major, minor = pynvml.nvmlDeviceGetCudaComputeCapability(handle)
|
||||
return DeviceCapability(major=major, minor=minor)
|
||||
except RuntimeError:
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=8)
|
||||
@with_nvml_context
|
||||
def has_device_capability(
|
||||
cls,
|
||||
capability: Union[Tuple[int, int], int],
|
||||
device_id: int = 0,
|
||||
) -> bool:
|
||||
try:
|
||||
return super().has_device_capability(capability, device_id)
|
||||
except RuntimeError:
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=8)
|
||||
@with_nvml_context
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
physical_device_id = device_id_to_physical_device_id(device_id)
|
||||
return cls._get_physical_device_name(physical_device_id)
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=8)
|
||||
@with_nvml_context
|
||||
def get_device_uuid(cls, device_id: int = 0) -> str:
|
||||
physical_device_id = device_id_to_physical_device_id(device_id)
|
||||
handle = pynvml.nvmlDeviceGetHandleByIndex(physical_device_id)
|
||||
return pynvml.nvmlDeviceGetUUID(handle)
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=8)
|
||||
@with_nvml_context
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
physical_device_id = device_id_to_physical_device_id(device_id)
|
||||
handle = pynvml.nvmlDeviceGetHandleByIndex(physical_device_id)
|
||||
return int(pynvml.nvmlDeviceGetMemoryInfo(handle).total)
|
||||
|
||||
@classmethod
|
||||
@with_nvml_context
|
||||
def is_full_nvlink(cls, physical_device_ids: List[int]) -> bool:
|
||||
"""
|
||||
query if the set of gpus are fully connected by nvlink (1 hop)
|
||||
"""
|
||||
handles = [
|
||||
pynvml.nvmlDeviceGetHandleByIndex(i) for i in physical_device_ids
|
||||
]
|
||||
for i, handle in enumerate(handles):
|
||||
for j, peer_handle in enumerate(handles):
|
||||
if i < j:
|
||||
try:
|
||||
p2p_status = pynvml.nvmlDeviceGetP2PStatus(
|
||||
handle,
|
||||
peer_handle,
|
||||
pynvml.NVML_P2P_CAPS_INDEX_NVLINK,
|
||||
)
|
||||
if p2p_status != pynvml.NVML_P2P_STATUS_OK:
|
||||
return False
|
||||
except pynvml.NVMLError:
|
||||
logger.exception(
|
||||
"NVLink detection failed. This is normal if"
|
||||
" your machine has no NVLink equipped.")
|
||||
return False
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def _get_physical_device_name(cls, device_id: int = 0) -> str:
|
||||
handle = pynvml.nvmlDeviceGetHandleByIndex(device_id)
|
||||
return pynvml.nvmlDeviceGetName(handle)
|
||||
|
||||
@classmethod
|
||||
@with_nvml_context
|
||||
def log_warnings(cls):
|
||||
device_ids: int = pynvml.nvmlDeviceGetCount()
|
||||
if device_ids > 1:
|
||||
device_names = [
|
||||
cls._get_physical_device_name(i) for i in range(device_ids)
|
||||
]
|
||||
if (len(set(device_names)) > 1
|
||||
and os.environ.get("CUDA_DEVICE_ORDER") != "PCI_BUS_ID"):
|
||||
logger.warning(
|
||||
"Detected different devices in the system: %s. Please"
|
||||
" make sure to set `CUDA_DEVICE_ORDER=PCI_BUS_ID` to "
|
||||
"avoid unexpected behavior.",
|
||||
", ".join(device_names),
|
||||
)
|
||||
|
||||
|
||||
class NonNvmlCudaPlatform(CudaPlatformBase):
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability:
|
||||
major, minor = torch.cuda.get_device_capability(device_id)
|
||||
return DeviceCapability(major=major, minor=minor)
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
return torch.cuda.get_device_name(device_id)
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
device_props = torch.cuda.get_device_properties(device_id)
|
||||
return device_props.total_memory
|
||||
|
||||
@classmethod
|
||||
def is_full_nvlink(cls, physical_device_ids: List[int]) -> bool:
|
||||
logger.exception(
|
||||
"NVLink detection not possible, as context support was"
|
||||
" not found. Assuming no NVLink available.")
|
||||
return False
|
||||
|
||||
|
||||
# Autodetect either NVML-enabled or non-NVML platform
|
||||
# based on whether NVML is available.
|
||||
nvml_available = False
|
||||
try:
|
||||
try:
|
||||
pynvml.nvmlInit()
|
||||
nvml_available = True
|
||||
except Exception:
|
||||
# On Jetson, NVML is not supported.
|
||||
nvml_available = False
|
||||
finally:
|
||||
if nvml_available:
|
||||
pynvml.nvmlShutdown()
|
||||
|
||||
CudaPlatform = NvmlCudaPlatform if nvml_available else NonNvmlCudaPlatform
|
||||
|
||||
try:
|
||||
from sphinx.ext.autodoc.mock import _MockModule
|
||||
|
||||
if not isinstance(pynvml, _MockModule):
|
||||
CudaPlatform.log_warnings()
|
||||
except ModuleNotFoundError:
|
||||
CudaPlatform.log_warnings()
|
||||
@@ -0,0 +1,194 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import enum
|
||||
import platform
|
||||
import random
|
||||
from platform import uname
|
||||
from typing import TYPE_CHECKING, NamedTuple, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class _Backend(enum.Enum):
|
||||
FLASH_ATTN = enum.auto()
|
||||
TORCH_SDPA = enum.auto()
|
||||
NO_ATTENTION = enum.auto()
|
||||
|
||||
|
||||
class PlatformEnum(enum.Enum):
|
||||
CUDA = enum.auto()
|
||||
OOT = enum.auto()
|
||||
UNSPECIFIED = enum.auto()
|
||||
|
||||
|
||||
class DeviceCapability(NamedTuple):
|
||||
major: int
|
||||
minor: int
|
||||
|
||||
def as_version_str(self) -> str:
|
||||
return f"{self.major}.{self.minor}"
|
||||
|
||||
def to_int(self) -> int:
|
||||
"""
|
||||
Express device capability as an integer ``<major><minor>``.
|
||||
|
||||
It is assumed that the minor version is always a single digit.
|
||||
"""
|
||||
assert 0 <= self.minor < 10
|
||||
return self.major * 10 + self.minor
|
||||
|
||||
|
||||
class Platform:
|
||||
_enum: PlatformEnum
|
||||
device_name: str
|
||||
device_type: str
|
||||
|
||||
# available dispatch keys:
|
||||
# check https://github.com/pytorch/pytorch/blob/313dac6c1ca0fa0cde32477509cce32089f8532a/torchgen/model.py#L134 # noqa
|
||||
# use "CPU" as a fallback for platforms not registered in PyTorch
|
||||
dispatch_key: str = "CPU"
|
||||
|
||||
supported_quantization: list[str] = []
|
||||
|
||||
def is_cuda(self) -> bool:
|
||||
return self._enum == PlatformEnum.CUDA
|
||||
|
||||
def is_out_of_tree(self) -> bool:
|
||||
return self._enum == PlatformEnum.OOT
|
||||
|
||||
def is_cuda_alike(self) -> bool:
|
||||
"""Stateless version of :func:`torch.cuda.is_available`."""
|
||||
return self._enum in (PlatformEnum.CUDA, PlatformEnum.ROCM)
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: _Backend, head_size: int,
|
||||
dtype: torch.dtype, kv_cache_dtype: Optional[str],
|
||||
block_size: int, use_v1: bool,
|
||||
use_mla: bool) -> str:
|
||||
"""Get the attention backend class of a device."""
|
||||
return ""
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(
|
||||
cls,
|
||||
device_id: int = 0,
|
||||
) -> Optional[DeviceCapability]:
|
||||
"""Stateless version of :func:`torch.cuda.get_device_capability`."""
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def has_device_capability(
|
||||
cls,
|
||||
capability: Union[Tuple[int, int], int],
|
||||
device_id: int = 0,
|
||||
) -> bool:
|
||||
"""
|
||||
Test whether this platform is compatible with a device capability.
|
||||
|
||||
The ``capability`` argument can either be:
|
||||
|
||||
- A tuple ``(major, minor)``.
|
||||
- An integer ``<major><minor>``. (See :meth:`DeviceCapability.to_int`)
|
||||
"""
|
||||
current_capability = cls.get_device_capability(device_id=device_id)
|
||||
if current_capability is None:
|
||||
return False
|
||||
|
||||
if isinstance(capability, tuple):
|
||||
return current_capability >= capability
|
||||
|
||||
return current_capability.to_int() >= capability
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
"""Get the name of a device."""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_uuid(cls, device_id: int = 0) -> str:
|
||||
"""Get the uuid of a device, e.g. the PCI bus ID."""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
"""Get the total memory of a device in bytes."""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: Optional[bool]) -> bool:
|
||||
"""
|
||||
Check if the current platform supports async output.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def inference_mode(cls):
|
||||
"""A device-specific wrapper of `torch.inference_mode`.
|
||||
|
||||
This wrapper is recommended because some hardware backends such as TPU
|
||||
do not support `torch.inference_mode`. In such a case, they will fall
|
||||
back to `torch.no_grad` by overriding this method.
|
||||
"""
|
||||
return torch.inference_mode(mode=True)
|
||||
|
||||
@classmethod
|
||||
def seed_everything(cls, seed: Optional[int] = None) -> None:
|
||||
"""
|
||||
Set the seed of each random module.
|
||||
`torch.manual_seed` will set seed on all devices.
|
||||
|
||||
Loosely based on: https://github.com/Lightning-AI/pytorch-lightning/blob/2.4.0/src/lightning/fabric/utilities/seed.py#L20
|
||||
"""
|
||||
if seed is not None:
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
|
||||
@classmethod
|
||||
def verify_model_arch(cls, model_arch: str) -> None:
|
||||
"""
|
||||
Verify whether the current platform supports the specified model
|
||||
architecture.
|
||||
|
||||
- This will raise an Error or Warning based on the model support on
|
||||
the current platform.
|
||||
- By default all models are considered supported.
|
||||
"""
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def verify_quantization(cls, quant: str) -> None:
|
||||
"""
|
||||
Verify whether the quantization is supported by the current platform.
|
||||
"""
|
||||
if cls.supported_quantization and \
|
||||
quant not in cls.supported_quantization:
|
||||
raise ValueError(
|
||||
f"{quant} quantization is currently not supported in "
|
||||
f"{cls.device_name}.")
|
||||
|
||||
@classmethod
|
||||
def get_current_memory_usage(cls,
|
||||
device: Optional[torch.types.Device] = None
|
||||
) -> float:
|
||||
"""
|
||||
Return the memory usage in bytes.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
"""
|
||||
Get device specific communicator class for distributed communication.
|
||||
"""
|
||||
return "fastvideo.v1.distributed.device_communicators.base_device_communicator.DeviceCommunicatorBase" # noqa
|
||||
|
||||
|
||||
class UnspecifiedPlatform(Platform):
|
||||
_enum = PlatformEnum.UNSPECIFIED
|
||||
device_type = ""
|
||||
@@ -0,0 +1,206 @@
|
||||
import os
|
||||
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
import sys
|
||||
# Fix the import path
|
||||
from fastvideo.v1.inference_engine import InferenceEngine
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.inference_args import prepare_inference_args
|
||||
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
|
||||
from fastvideo.v1.logger import logger
|
||||
from fastvideo.v1.models.loader.fsdp_load import get_param_names_mapping
|
||||
from safetensors.torch import safe_open, save_file
|
||||
import shutil
|
||||
|
||||
def initialize_distributed_and_parallelism(inference_args: InferenceArgs):
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
torch.cuda.set_device(local_rank)
|
||||
init_distributed_environment(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank
|
||||
)
|
||||
device_str = f"cuda:{local_rank}"
|
||||
inference_args.device_str = device_str
|
||||
inference_args.device = torch.device(device_str)
|
||||
initialize_model_parallel(
|
||||
sequence_model_parallel_size=inference_args.sp_size,
|
||||
tensor_model_parallel_size=inference_args.tp_size,
|
||||
)
|
||||
|
||||
def convert_and_save_lora_weights(input_folder, output_folder, lora_param_mapping_fn):
|
||||
|
||||
# Check if the input folder exists
|
||||
if not os.path.exists(input_folder):
|
||||
raise ValueError(f"Input folder {input_folder} not found")
|
||||
|
||||
# Create output folder if it doesn't exist
|
||||
os.makedirs(output_folder, exist_ok=True)
|
||||
|
||||
# Process safetensors file
|
||||
safetensors_file = os.path.join(input_folder, "pytorch_lora_weights.safetensors")
|
||||
if os.path.exists(safetensors_file):
|
||||
# Load all tensors from the safetensors file
|
||||
tensors = {}
|
||||
with safe_open(safetensors_file, framework="pt", device="cpu") as f:
|
||||
tensor_names = f.keys()
|
||||
print(f"Found {len(tensor_names)} tensors in {safetensors_file}")
|
||||
|
||||
for name in tensor_names:
|
||||
tensors[name] = f.get_tensor(name)
|
||||
|
||||
# Extract base names (without LoRA suffix) and LoRA suffixes
|
||||
base_names_dict = {}
|
||||
lora_suffixes = {}
|
||||
|
||||
for name in tensor_names:
|
||||
if name.endswith(".lora_A.weight") or name.endswith(".lora_B.weight"):
|
||||
for suffix in [".lora_A.weight", ".lora_B.weight"]:
|
||||
if name.endswith(suffix):
|
||||
base_name = name[:-len(suffix)] + ".weight"
|
||||
base_names_dict[name] = base_name
|
||||
lora_suffixes[name] = suffix
|
||||
break
|
||||
else:
|
||||
base_names_dict[name] = name
|
||||
lora_suffixes[name] = ""
|
||||
|
||||
# Apply the mapping function to get all mappings at once
|
||||
name_mappings = {}
|
||||
for full_name, base_name in base_names_dict.items():
|
||||
try:
|
||||
# First remove "transformer." prefix if it exists
|
||||
has_transformer_prefix = False
|
||||
processed_base_name = base_name
|
||||
if base_name.startswith("transformer."):
|
||||
processed_base_name = base_name[12:] # Remove "transformer." (12 characters)
|
||||
has_transformer_prefix = True
|
||||
|
||||
# Apply the mapping function to the processed name
|
||||
target_base_name, merge_index, total_splitted_params = lora_param_mapping_fn(processed_base_name)
|
||||
|
||||
# Add back the "transformer." prefix if it was removed
|
||||
if target_base_name is not None and has_transformer_prefix:
|
||||
target_base_name = "transformer." + target_base_name
|
||||
|
||||
if target_base_name is not None:
|
||||
name_mappings[full_name] = (target_base_name, merge_index, total_splitted_params)
|
||||
except Exception as e:
|
||||
print(f"Error mapping parameter {base_name}: {e}")
|
||||
|
||||
# Create the converted weights using the mappings
|
||||
converted_weights = {}
|
||||
for name, tensor in tensors.items():
|
||||
suffix = lora_suffixes.get(name, "")
|
||||
|
||||
if name in name_mappings:
|
||||
target_base_name, merge_index, _ = name_mappings[name]
|
||||
|
||||
# Reconstruct full parameter name with LoRA suffix
|
||||
if target_base_name.endswith('.weight'):
|
||||
target_base_name_without_suffix = target_base_name[:-7] # 移除 .weight
|
||||
new_key = target_base_name_without_suffix + suffix
|
||||
else:
|
||||
new_key = target_base_name + suffix
|
||||
|
||||
if merge_index is not None:
|
||||
new_key += f"_{merge_index}"
|
||||
|
||||
converted_weights[new_key] = tensor
|
||||
else:
|
||||
# Keep original name if no mapping was found
|
||||
converted_weights[name] = tensor
|
||||
|
||||
print(f"Converted {len(converted_weights)} parameters")
|
||||
|
||||
# Save the converted weights
|
||||
output_safetensors_file = os.path.join(output_folder, "pytorch_lora_weights.safetensors")
|
||||
save_file(converted_weights, output_safetensors_file)
|
||||
print(f"Converted weights saved to {output_safetensors_file}")
|
||||
else:
|
||||
print(f"Warning: No safetensors file found in {input_folder}")
|
||||
|
||||
# Copy other files (like config and optimizer state) without modification
|
||||
for filename in os.listdir(input_folder):
|
||||
if filename != "pytorch_lora_weights.safetensors":
|
||||
input_path = os.path.join(input_folder, filename)
|
||||
output_path = os.path.join(output_folder, filename)
|
||||
|
||||
if os.path.isfile(input_path):
|
||||
shutil.copy2(input_path, output_path)
|
||||
print(f"Copied {filename} to output folder")
|
||||
|
||||
print(f"LoRA conversion complete: {input_folder} → {output_folder}")
|
||||
|
||||
def main(inference_args: InferenceArgs):
|
||||
initialize_distributed_and_parallelism(inference_args)
|
||||
engine = InferenceEngine.create_engine(
|
||||
inference_args,
|
||||
)
|
||||
|
||||
if inference_args.prompt_path is not None:
|
||||
with open(inference_args.prompt_path) as f:
|
||||
prompts = [line.strip() for line in f.readlines()]
|
||||
else:
|
||||
prompts = [inference_args.prompt]
|
||||
# from IPython import embed; embed()
|
||||
|
||||
# # convert lora format
|
||||
# input_path = "data/Hunyuan-Black-Myth-Wukong-lora-weight/"
|
||||
# output_path = "data/Hunyuan-Black-Myth-Wukong-lora-weight_converted"
|
||||
# # from IPython import embed; embed()
|
||||
# lora_param_mapping_fn = get_param_names_mapping(engine.pipeline.transformer._param_names_mapping)
|
||||
# converted_weights = convert_and_save_lora_weights(
|
||||
# input_path,
|
||||
# output_path,
|
||||
# lora_param_mapping_fn,
|
||||
# )
|
||||
|
||||
# Process each prompt
|
||||
for prompt in prompts:
|
||||
# lora_checkpoint = "data/Hunyuan-Black-Myth-Wukong-lora-weight_converted/"
|
||||
# if lora_checkpoint:
|
||||
# import json
|
||||
# print(f"Loading LoRA weights from lora_checkpoint: {lora_checkpoint}")
|
||||
# config_path = os.path.join(lora_checkpoint, "lora_config.json")
|
||||
# with open(config_path, "r") as f:
|
||||
# lora_config_dict = json.load(f)
|
||||
# rank = lora_config_dict["lora_params"]["lora_rank"]
|
||||
# lora_alpha = lora_config_dict["lora_params"]["lora_alpha"]
|
||||
# lora_scaling = lora_alpha / rank
|
||||
# engine.pipeline.load_lora_weights(lora_checkpoint, adapter_name="default")
|
||||
# from IPython import embed; embed()
|
||||
# engine.pipeline.set_adapters(["default"], [lora_scaling])
|
||||
# print(f"Successfully Loaded LoRA weights from {lora_checkpoint}")
|
||||
outputs = engine.run(
|
||||
prompt=prompt,
|
||||
inference_args=inference_args,
|
||||
)
|
||||
# img_attn_qkv, img_attn_proj, linear1, linear2
|
||||
# Process outputs
|
||||
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
|
||||
# Save video
|
||||
os.makedirs(os.path.dirname(inference_args.output_path), exist_ok=True)
|
||||
imageio.mimsave(
|
||||
os.path.join(inference_args.output_path, f"{prompt[:100]}.mp4"),
|
||||
frames,
|
||||
fps=inference_args.fps
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
inference_args = prepare_inference_args(sys.argv[1:])
|
||||
main(inference_args)
|
||||
@@ -0,0 +1,18 @@
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
import os
|
||||
|
||||
def main():
|
||||
print(os.environ["RANK"])
|
||||
print(os.environ["WORLD_SIZE"])
|
||||
print(os.environ["LOCAL_RANK"])
|
||||
print(os.environ["MASTER_ADDR"])
|
||||
print(os.environ["MASTER_PORT"])
|
||||
print(os.environ["RANK"])
|
||||
print(os.environ["WORLD_SIZE"])
|
||||
print(os.environ["LOCAL_RANK"])
|
||||
print(os.environ["MASTER_ADDR"])
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,19 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=8
|
||||
torchrun --standalone --nnodes=1 --nproc_per_node=$num_gpus \
|
||||
--master_port 29503 \
|
||||
tp_example.py
|
||||
|
||||
|
||||
num_gpus=2
|
||||
torchrun --standalone --nnodes=4 --nproc_per_node=$num_gpus \
|
||||
--master_port 29503 \
|
||||
test_hunyuanvideo.py --sequence_model_parallel_size $num_gpus
|
||||
|
||||
|
||||
|
||||
num_gpus=2
|
||||
torchrun --standalone --nnodes=1 --nproc_per_node=$num_gpus \
|
||||
--master_port 29503 \
|
||||
fastvideo/v1/tests/test_hunyuanvideo_load.py --sequence_model_parallel_size $num_gpus
|
||||
@@ -0,0 +1,207 @@
|
||||
from fastvideo.v1.models.encoders.clip import CLIPTextModel
|
||||
# TODO: check if correct
|
||||
from fastvideo.models.hunyuan.text_encoder import TextEncoder, load_text_encoder, load_tokenizer
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import argparse
|
||||
import numpy as np
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from transformers import AutoConfig
|
||||
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
|
||||
# from fastvideo.v1.models.hunyuan.text_encoder import load_text_encoder, load_tokenizer
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
def initialize_identical_weights(model1, model2, seed=42):
|
||||
"""Initialize both models with identical weights using a fixed seed for reproducibility."""
|
||||
# Get all parameters from both models
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
# Initialize each layer with identical values
|
||||
with torch.no_grad():
|
||||
# Initialize weights
|
||||
for name1, param1 in params1.items():
|
||||
if 'weight' in name1:
|
||||
# Set seed before each weight initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'weight' in name2:
|
||||
# Reset seed to get same initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
|
||||
# Initialize biases
|
||||
for name1, param1 in params1.items():
|
||||
if 'bias' in name1:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
param1.data = param1.data.to(torch.bfloat16)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'bias' in name2:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
param2.data = param2.data.to(torch.bfloat16)
|
||||
|
||||
logger.info("Both models initialized with identical weights in bfloat16")
|
||||
return model1, model2
|
||||
|
||||
def setup_args():
|
||||
parser = argparse.ArgumentParser(description='CLIP Text Encoder Test')
|
||||
parser.add_argument('--model-path', type=str, default="openai/clip-vit-large-patch14",
|
||||
help='Path to the CLIP model')
|
||||
parser.add_argument('--precision', type=str, default="float16",
|
||||
help='Precision to use for the model (float32, float16, bfloat16)')
|
||||
return parser.parse_args()
|
||||
|
||||
def test_clip_encoder():
|
||||
init_distributed_environment(world_size=1, rank=0, distributed_init_method="env://", local_rank=0, backend="nccl")
|
||||
initialize_model_parallel(tensor_model_parallel_size=1, sequence_model_parallel_size=1, backend="nccl")
|
||||
args = setup_args()
|
||||
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
# Initialize the two model implementations
|
||||
logger.info(f"Loading models from {args.model_path}")
|
||||
model_path = "data/hunyuanvideo-community/HunyuanVideo/text_encoder_2"
|
||||
|
||||
# config = json.load(open(os.path.join(model_path, "config.json")))
|
||||
|
||||
hf_config = AutoConfig.from_pretrained(model_path)
|
||||
print(hf_config)
|
||||
print(hf_config.use_return_dict)
|
||||
|
||||
# Load our implementation using the loader from text_encoder/__init__.py
|
||||
model1, _ = load_text_encoder(
|
||||
text_encoder_type="clipL",
|
||||
text_encoder_precision='fp16',
|
||||
text_encoder_path=model_path,
|
||||
logger=logger,
|
||||
device=device
|
||||
)
|
||||
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
loader = TextEncoderLoader()
|
||||
args.device_str = "cuda:0"
|
||||
model2 = loader.load_model(model_path, hf_config, device)
|
||||
|
||||
# Load the HuggingFace implementation directly
|
||||
# model2 = CLIPTextModel(hf_config)
|
||||
model2 = model2.to(torch.float16)
|
||||
model2 = model2.to(device)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights between the two models
|
||||
logger.info("Comparing model weights for sanity check...")
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
# Check number of parameters
|
||||
logger.info(f"Model1 has {len(params1)} parameters")
|
||||
logger.info(f"Model2 has {len(params2)} parameters")
|
||||
|
||||
# Compare a few key parameters
|
||||
|
||||
# weight_diffs = []
|
||||
# for (name1, param1), (name2, param2) in zip(
|
||||
# sorted(params1.items()), sorted(params2.items())
|
||||
# ):
|
||||
# # if len(weight_diffs) < 5: # Just check a few parameters
|
||||
# max_diff = torch.max(torch.abs(param1 - param2)).item()
|
||||
# mean_diff = torch.mean(torch.abs(param1 - param2)).item()
|
||||
# weight_diffs.append((name1, name2, max_diff, mean_diff))
|
||||
# logger.info(f"Parameter: {name1} vs {name2}")
|
||||
# logger.info(f" Max diff: {max_diff}, Mean diff: {mean_diff}")
|
||||
|
||||
# Load tokenizer
|
||||
tokenizer, _ = load_tokenizer(
|
||||
tokenizer_type="clipL",
|
||||
tokenizer_path=args.model_path,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
# Test with some sample prompts
|
||||
prompts = [
|
||||
"a photo of a cat",
|
||||
"a beautiful landscape with mountains",
|
||||
"an astronaut riding a horse on the moon"
|
||||
]
|
||||
|
||||
logger.info("Testing CLIP text encoder with sample prompts")
|
||||
|
||||
with torch.no_grad():
|
||||
for prompt in prompts:
|
||||
logger.info(f"Testing prompt: '{prompt}'")
|
||||
|
||||
# Tokenize the prompt
|
||||
tokens = tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=77,
|
||||
truncation=True,
|
||||
return_tensors="pt"
|
||||
).to(device)
|
||||
|
||||
# Get embeddings from our implementation
|
||||
outputs1 = model1(
|
||||
input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
output_hidden_states=True
|
||||
)
|
||||
|
||||
logger.info(f"Testing model2")
|
||||
|
||||
# Get embeddings from HuggingFace implementation
|
||||
outputs2 = model2(
|
||||
input_ids=tokens.input_ids,
|
||||
# attention_mask=tokens.attention_mask,
|
||||
output_hidden_states=True
|
||||
)
|
||||
|
||||
# Compare last hidden states
|
||||
last_hidden_state1 = outputs1.last_hidden_state[tokens.attention_mask==1]
|
||||
last_hidden_state2 = outputs2.last_hidden_state[tokens.attention_mask==1]
|
||||
# print("last_hidden_state1", last_hidden_state1)
|
||||
# print("last_hidden_state2", last_hidden_state2)
|
||||
|
||||
assert last_hidden_state1.shape == last_hidden_state2.shape, \
|
||||
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
|
||||
|
||||
max_diff_hidden = torch.max(torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
mean_diff_hidden = torch.mean(torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
|
||||
logger.info(f"Maximum difference in last hidden states: {max_diff_hidden.item()}")
|
||||
logger.info(f"Mean difference in last hidden states: {mean_diff_hidden.item()}")
|
||||
|
||||
# Compare pooler outputs
|
||||
pooler_output1 = outputs1.pooler_output
|
||||
pooler_output2 = outputs2.pooler_output
|
||||
|
||||
assert pooler_output1.shape == pooler_output2.shape, \
|
||||
f"Pooler output shapes don't match: {pooler_output1.shape} vs {pooler_output2.shape}"
|
||||
|
||||
max_diff_pooler = torch.max(torch.abs(pooler_output1 - pooler_output2))
|
||||
mean_diff_pooler = torch.mean(torch.abs(pooler_output1 - pooler_output2))
|
||||
|
||||
logger.info(f"Maximum difference in pooler outputs: {max_diff_pooler.item()}")
|
||||
logger.info(f"Mean difference in pooler outputs: {mean_diff_pooler.item()}")
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
assert max_diff_hidden < 1e-4, \
|
||||
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
|
||||
assert max_diff_pooler < 1e-4, \
|
||||
f"Pooler outputs differ significantly: max diff = {max_diff_pooler.item()}"
|
||||
|
||||
logger.info("Test passed! Both CLIP text encoder implementations produce similar outputs.")
|
||||
logger.info("Test completed successfully")
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_clip_encoder()
|
||||
@@ -0,0 +1,161 @@
|
||||
from fastvideo.v1.models.vaes.hunyuanvae import AutoencoderKLHunyuanVideo as MyHunyuanVAE
|
||||
from diffusers import AutoencoderKLHunyuanVideo as DiffusersHunyuanVAE
|
||||
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import argparse
|
||||
import numpy as np
|
||||
import json
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from safetensors.torch import load_file
|
||||
logger = init_logger(__name__)
|
||||
|
||||
def initialize_identical_weights(model1, model2, seed=42):
|
||||
"""Initialize both models with identical weights using a fixed seed for reproducibility."""
|
||||
# Get all parameters from both models
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
# Initialize each layer with identical values
|
||||
with torch.no_grad():
|
||||
# Initialize weights
|
||||
for name1, param1 in params1.items():
|
||||
if 'weight' in name1:
|
||||
# Set seed before each weight initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'weight' in name2:
|
||||
# Reset seed to get same initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
|
||||
# Initialize biases
|
||||
for name1, param1 in params1.items():
|
||||
if 'bias' in name1:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
param1.data = param1.data.to(torch.bfloat16)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'bias' in name2:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
param2.data = param2.data.to(torch.bfloat16)
|
||||
|
||||
logger.info("Both models initialized with identical weights in bfloat16")
|
||||
return model1, model2
|
||||
|
||||
def setup_args():
|
||||
parser = argparse.ArgumentParser(description='HunyuanVAE Test')
|
||||
parser.add_argument('--in-channels', type=int, default=4,
|
||||
help='Number of input channels')
|
||||
parser.add_argument('--out-channels', type=int, default=4,
|
||||
help='Number of output channels')
|
||||
parser.add_argument('--latent-channels', type=int, default=4,
|
||||
help='Number of latent channels')
|
||||
return parser.parse_args()
|
||||
|
||||
def test_hunyuan_vae():
|
||||
args = setup_args()
|
||||
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
# Model parameters
|
||||
in_channels = args.in_channels
|
||||
out_channels = args.out_channels
|
||||
latent_channels = args.latent_channels
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
print
|
||||
# Initialize the two model implementations
|
||||
path = "data/hunyuanvideo-community/HunyuanVideo/vae"
|
||||
config_path = os.path.join(path, "config.json")
|
||||
config = json.load(open(config_path))
|
||||
config.pop("_class_name")
|
||||
config.pop("_diffusers_version")
|
||||
model1 = MyHunyuanVAE(
|
||||
**config
|
||||
).to(torch.bfloat16)
|
||||
|
||||
model2 = DiffusersHunyuanVAE(
|
||||
**config
|
||||
).to(torch.bfloat16)
|
||||
|
||||
loaded = load_file(os.path.join(path, "diffusion_pytorch_model.safetensors"))
|
||||
model1.load_state_dict(loaded)
|
||||
model2.load_state_dict(loaded)
|
||||
|
||||
|
||||
# Set both models to eval mode
|
||||
model1.eval()
|
||||
model2.eval()
|
||||
|
||||
# Move to GPU
|
||||
model1 = model1.to(device)
|
||||
model2 = model2.to(device)
|
||||
|
||||
|
||||
model1.enable_tiling(
|
||||
tile_sample_min_height=32,
|
||||
tile_sample_min_width=32,
|
||||
tile_sample_min_num_frames=8,
|
||||
tile_sample_stride_height=16,
|
||||
tile_sample_stride_width=16,
|
||||
tile_sample_stride_num_frames=4
|
||||
)
|
||||
model2.enable_tiling(
|
||||
tile_sample_min_height=32,
|
||||
tile_sample_min_width=32,
|
||||
tile_sample_min_num_frames=8,
|
||||
tile_sample_stride_height=16,
|
||||
tile_sample_stride_width=16,
|
||||
tile_sample_stride_num_frames=4
|
||||
)
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
|
||||
# Video input [B, C, T, H, W]
|
||||
input_tensor = torch.randn(batch_size, 3, 21, 64, 64, device=device, dtype=torch.bfloat16)
|
||||
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
# Test encoding
|
||||
logger.info("Testing encoding...")
|
||||
latent1 = model1.encode(input_tensor).mean
|
||||
print("--------------------------------")
|
||||
latent2 = model2.encode(input_tensor).latent_dist.mean
|
||||
# Check if latents have the same shape
|
||||
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
|
||||
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
|
||||
# Check if latents are similar
|
||||
max_diff_encode = torch.max(torch.abs(latent1 - latent2))
|
||||
mean_diff_encode = torch.mean(torch.abs(latent1 - latent2))
|
||||
logger.info(f"Maximum difference between encoded latents: {max_diff_encode.item()}")
|
||||
logger.info(f"Mean difference between encoded latents: {mean_diff_encode.item()}")
|
||||
assert max_diff_encode < 1e-4, f"Encoded latents differ significantly: max diff = {max_diff_encode.item()}"
|
||||
# Test decoding
|
||||
logger.info("Testing decoding...")
|
||||
output1 = model1.decode(latent1)
|
||||
output2 = model2.decode(latent2).sample
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
|
||||
# Check if outputs are similar
|
||||
max_diff_decode = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff_decode = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info(f"Maximum difference between decoded outputs: {max_diff_decode.item()}")
|
||||
logger.info(f"Mean difference between decoded outputs: {mean_diff_decode.item()}")
|
||||
assert max_diff_decode < 1e-4, f"Decoded outputs differ significantly: max diff = {max_diff_decode.item()}"
|
||||
|
||||
|
||||
logger.info("Test passed! Both VAE implementations produce similar outputs.")
|
||||
logger.info("Test completed successfully")
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_hunyuan_vae()
|
||||
@@ -0,0 +1,233 @@
|
||||
from itertools import chain
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import argparse
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size,
|
||||
destroy_model_parallel,
|
||||
destroy_distributed_environment,
|
||||
cleanup_dist_env_and_memory
|
||||
)
|
||||
|
||||
from fastvideo.v1.utils.parallel_states import initialize_sequence_parallel_state
|
||||
from torch.distributed._composable.fsdp import CPUOffloadPolicy, fully_shard
|
||||
from torch.distributed.device_mesh import init_device_mesh
|
||||
from fastvideo.v1.models.loader.fsdp_load import shard_model
|
||||
from fastvideo.v1.models.dits.hunyuanvideo import HunyuanVideoTransformer3DModel as HunyuanVideoDit
|
||||
from fastvideo.v1.models.hunyuan.modules.models import HYVideoDiffusionTransformer
|
||||
from fastvideo.v1.models.hunyuan_hf.modeling_hunyuan import HunyuanVideoTransformer3DModel
|
||||
logger = init_logger(__name__)
|
||||
|
||||
def initialize_identical_weights(model1, model2, seed=42):
|
||||
"""Initialize both models with identical weights using a fixed seed for reproducibility."""
|
||||
# Get all parameters from both models
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
# Initialize each layer with identical values
|
||||
with torch.no_grad():
|
||||
# Initialize weights
|
||||
for name1, param1 in params1.items():
|
||||
if 'weight' in name1:
|
||||
# Set seed before each weight initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'weight' in name2:
|
||||
# Reset seed to get same initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
|
||||
# Initialize biases
|
||||
for name1, param1 in params1.items():
|
||||
if 'bias' in name1:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
param1.data = param1.data.to(torch.bfloat16)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'bias' in name2:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
param2.data = param2.data.to(torch.bfloat16)
|
||||
|
||||
logger.info("Both models initialized with identical weights in bfloat16")
|
||||
return model1, model2
|
||||
|
||||
def setup_args():
|
||||
parser = argparse.ArgumentParser(description='Distributed HunyuanVideo Test')
|
||||
parser.add_argument('--sequence_model_parallel_size', type=int, default=1,
|
||||
help='Degree of sequence model parallelism')
|
||||
parser.add_argument('--hidden-size', type=int, default=128,
|
||||
help='Hidden size for the model')
|
||||
parser.add_argument('--heads-num', type=int, default=4,
|
||||
help='Number of attention heads')
|
||||
parser.add_argument('--double-blocks-depth', type=int, default=2,
|
||||
help='Number of double stream blocks')
|
||||
parser.add_argument('--single-blocks-depth', type=int, default=2,
|
||||
help='Number of single stream blocks')
|
||||
return parser.parse_args()
|
||||
|
||||
def test_hunyuanvideo_distributed():
|
||||
args = setup_args()
|
||||
|
||||
# Initialize distributed environment
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
|
||||
logger.info(f"Initializing process: rank={rank}, local_rank={local_rank}, world_size={world_size}")
|
||||
|
||||
# Initialize distributed environment
|
||||
init_distributed_environment(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank
|
||||
)
|
||||
|
||||
# Initialize tensor model parallel groups
|
||||
initialize_model_parallel(
|
||||
sequence_model_parallel_size=args.sequence_model_parallel_size
|
||||
)
|
||||
initialize_sequence_parallel_state(world_size)
|
||||
# Get tensor parallel info
|
||||
sp_rank = get_sequence_model_parallel_rank()
|
||||
sp_world_size = get_sequence_model_parallel_world_size()
|
||||
|
||||
logger.info(f"Process rank {rank} initialized with SP rank {sp_rank} in SP world size {sp_world_size}")
|
||||
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
# Small model parameters for testing
|
||||
hidden_size = args.hidden_size
|
||||
heads_num = args.heads_num
|
||||
mm_double_blocks_depth = args.double_blocks_depth
|
||||
mm_single_blocks_depth = args.single_blocks_depth
|
||||
patch_size = [1, 2, 2]
|
||||
torch.cuda.set_device(f"cuda:{local_rank}")
|
||||
# Initialize the two model implementations
|
||||
model1 = HunyuanVideoDit(
|
||||
patch_size=2,
|
||||
patch_size_t=1,
|
||||
in_channels=4,
|
||||
out_channels=4,
|
||||
attention_head_dim=hidden_size//heads_num,
|
||||
num_attention_heads=heads_num,
|
||||
num_layers=mm_double_blocks_depth,
|
||||
num_single_layers=mm_single_blocks_depth,
|
||||
rope_axes_dim=[8, 16, 8], # sum = hidden_size // heads_num = 32
|
||||
dtype=torch.bfloat16
|
||||
).to(torch.bfloat16)
|
||||
model2 = HYVideoDiffusionTransformer(
|
||||
patch_size=patch_size,
|
||||
in_channels=4,
|
||||
hidden_size=hidden_size,
|
||||
heads_num=heads_num,
|
||||
mm_double_blocks_depth=mm_double_blocks_depth,
|
||||
mm_single_blocks_depth=mm_single_blocks_depth,
|
||||
rope_dim_list=[8, 16, 8], # sum = hidden_size // heads_num = 32
|
||||
dtype=torch.bfloat16
|
||||
).to(torch.bfloat16)
|
||||
|
||||
|
||||
# print("--------------------------------")
|
||||
# for name, param in model3.named_parameters():
|
||||
# print(name)
|
||||
# import pdb; pdb.set_trace()
|
||||
# # Initialize with identical weights
|
||||
model1, model2 = initialize_identical_weights(model1, model2, seed=42)
|
||||
device_mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(sp_world_size,),
|
||||
mesh_dim_names=("dp", ),
|
||||
)
|
||||
shard_model(model1, cpu_offload=False, reshard_after_forward=True)
|
||||
for n, p in chain(model1.named_parameters(), model1.named_buffers()):
|
||||
if p.is_meta:
|
||||
raise RuntimeError(f"Unexpected param or buffer {n} on meta device.")
|
||||
for p in model1.parameters():
|
||||
p.requires_grad = False
|
||||
# Set both models to eval mode
|
||||
model1.eval()
|
||||
model2.eval()
|
||||
|
||||
# Move to GPU based on local rank (0 or 1 for 2 GPUs)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
model1 = model1.to(device)
|
||||
model2 = model2.to(device)
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
seq_len = 3
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size, 4, 8, 16, 16, device=device, dtype=torch.bfloat16)
|
||||
chunk_per_rank = hidden_states.shape[2] // sp_world_size
|
||||
hidden_states = hidden_states[:, :, sp_rank * chunk_per_rank:(sp_rank + 1) * chunk_per_rank]
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size, seq_len + 1, 4096, device=device, dtype=torch.bfloat16)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([500], device=device, dtype=torch.bfloat16)
|
||||
|
||||
# Attention mask for text
|
||||
encoder_attention_mask = torch.ones(batch_size, seq_len, device=device, dtype=torch.bfloat16)
|
||||
guidance = torch.tensor([1.0], device=device, dtype=torch.bfloat16)
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
output1 = model1(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
)
|
||||
print("--------------------------------")
|
||||
output2, _ = model2(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
logger.info(f"Maximum difference between outputs: {max_diff.item()}")
|
||||
# mean diff
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info(f"Mean difference between outputs: {mean_diff.item()}")
|
||||
# diff sum
|
||||
diff_sum = torch.sum(torch.abs(output1 - output2))
|
||||
logger.info(f"Diff sum between outputs: {diff_sum.item()}")
|
||||
# sum
|
||||
sum_output1 = torch.sum(output1.float())
|
||||
sum_output2 = torch.sum(output2.float())
|
||||
logger.info(f"Rank {sp_rank} Sum of output1: {sum_output1.item()}")
|
||||
logger.info(f"Rank {sp_rank} Sum of output2: {sum_output2.item()}")
|
||||
# The outputs should be very close if not identical
|
||||
assert max_diff < 1e-3, f"Outputs differ significantly: max diff = {max_diff.item()}" # Increased tolerance for bf16
|
||||
|
||||
logger.info("Test passed! Both model implementations produce the same outputs.")
|
||||
|
||||
# Clean up
|
||||
logger.info("Cleaning up distributed environment")
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
logger.info("Test completed successfully")
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_hunyuanvideo_distributed()
|
||||
@@ -0,0 +1,195 @@
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import argparse
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size,
|
||||
destroy_model_parallel,
|
||||
destroy_distributed_environment,
|
||||
cleanup_dist_env_and_memory
|
||||
)
|
||||
import json
|
||||
from fastvideo.v1.models.dits.hunyuanvideo import HunyuanVideoTransformer3DModel as HunyuanVideoDit
|
||||
from fastvideo.models.hunyuan.modules.models import HUNYUAN_VIDEO_CONFIG
|
||||
from fastvideo.models.hunyuan.modules.models import HYVideoDiffusionTransformer
|
||||
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
|
||||
from fastvideo.models.hunyuan.inference import Inference
|
||||
import glob
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def setup_args():
|
||||
parser = argparse.ArgumentParser(description='Distributed HunyuanVideo Test')
|
||||
parser.add_argument('--sequence_model_parallel_size', type=int, default=1,
|
||||
help='Degree of sequence model parallelism')
|
||||
return parser.parse_args()
|
||||
|
||||
def test_hunyuanvideo_distributed():
|
||||
args = setup_args()
|
||||
|
||||
# Initialize distributed environment
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
|
||||
logger.info(f"Initializing process: rank={rank}, local_rank={local_rank}, world_size={world_size}")
|
||||
|
||||
# Initialize distributed environment
|
||||
init_distributed_environment(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank
|
||||
)
|
||||
torch.cuda.set_device(f"cuda:{local_rank}")
|
||||
# Initialize tensor model parallel groups
|
||||
initialize_model_parallel(
|
||||
sequence_model_parallel_size=args.sequence_model_parallel_size
|
||||
)
|
||||
|
||||
# Get tensor parallel info
|
||||
sp_rank = get_sequence_model_parallel_rank()
|
||||
sp_world_size = get_sequence_model_parallel_world_size()
|
||||
|
||||
logger.info(f"Process rank {rank} initialized with SP rank {sp_rank} in SP world size {sp_world_size}")
|
||||
|
||||
# load data/hunyuanvideo_community/transformer/config.json
|
||||
with open("data/hunyuanvideo-community/HunyuanVideo/transformer/config.json", "r") as f:
|
||||
config = json.load(f)
|
||||
# remove "_class_name": "HunyuanVideoTransformer3DModel", "_diffusers_version": "0.32.0.dev0",
|
||||
# TODO: write normalize config function
|
||||
config.pop("_class_name")
|
||||
config.pop("_diffusers_version")
|
||||
# load data/hunyuanvideo_community/transformer/*.safetensors
|
||||
weight_dir_list = glob.glob("data/hunyuanvideo-community/HunyuanVideo/transformer/*.safetensors")
|
||||
# to str
|
||||
weight_dir_list = [str(path) for path in weight_dir_list]
|
||||
model1 = load_fsdp_model(
|
||||
HunyuanVideoDit,
|
||||
init_params=config,
|
||||
weight_dir_list=weight_dir_list,
|
||||
device=torch.device(f"cuda:{local_rank}"),
|
||||
cpu_offload=False
|
||||
)
|
||||
|
||||
# successfully sharded the model (hunyuanvideo bf16 should take around 26GB in total)
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
logger.info(f"Total parameters: {total_params / 1e9}B")
|
||||
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
|
||||
model2 = HYVideoDiffusionTransformer(
|
||||
in_channels=16,
|
||||
out_channels=16,
|
||||
**HUNYUAN_VIDEO_CONFIG["HYVideo-T/2-cfgdistill"],
|
||||
device=torch.device(f"cuda:{local_rank}"),
|
||||
dtype=torch.bfloat16
|
||||
).bfloat16()
|
||||
# data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt
|
||||
state_dict = torch.load("/mbz/users/hao.zhang/peiyuan/FastVideo/data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt", map_location=lambda storage, loc: storage)["module"]
|
||||
model2.load_state_dict(state_dict, strict=True)
|
||||
model2.to(torch.device(f"cuda:{local_rank}")).bfloat16()
|
||||
print("load state dict done")
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info(f"Model 2 weight sum: {weight_sum_model2}")
|
||||
logger.info(f"Model 2 weight mean: {weight_mean_model2}")
|
||||
|
||||
# Set both models to eval mode
|
||||
model1.eval()
|
||||
model2.eval()
|
||||
|
||||
# Create random inputs for testing
|
||||
batch_size = 1
|
||||
seq_len = 3
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size, 16, 8, 16, 16, device=device, dtype=torch.bfloat16)
|
||||
chunk_per_rank = hidden_states.shape[2] // sp_world_size
|
||||
hidden_states = hidden_states[:, :, sp_rank * chunk_per_rank:(sp_rank + 1) * chunk_per_rank]
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size, seq_len + 1, 4096, device=device, dtype=torch.bfloat16)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([500], device=device, dtype=torch.bfloat16)
|
||||
|
||||
# Attention mask for text
|
||||
encoder_attention_mask = torch.ones(batch_size, seq_len, device=device, dtype=torch.bfloat16)
|
||||
guidance = torch.tensor([1.0], device=device, dtype=torch.bfloat16)
|
||||
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
# Run inference on model1
|
||||
logger.info(f"Running inference on model1")
|
||||
output1 = model1(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
)
|
||||
logger.info("Model 1 inference completed")
|
||||
|
||||
# Run inference on model2
|
||||
output2, _ = model2(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
)
|
||||
logger.info("Model 2 inference completed")
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
|
||||
# Compare weight sums and means
|
||||
logger.info(f"Model 1 weight sum: {weight_sum_model1}")
|
||||
logger.info(f"Model 2 weight sum: {weight_sum_model2}")
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info(f"Weight sum difference: {weight_sum_diff}")
|
||||
|
||||
logger.info(f"Model 1 weight mean: {weight_mean_model1}")
|
||||
logger.info(f"Model 2 weight mean: {weight_mean_model2}")
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info(f"Weight mean difference: {weight_mean_diff}")
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
assert max_diff < 1e-2, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
|
||||
# mean diff
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
|
||||
# diff sum
|
||||
diff_sum = torch.sum(torch.abs(output1 - output2))
|
||||
logger.info(f"Diff sum between outputs: {diff_sum.item()}")
|
||||
|
||||
# sum
|
||||
sum_output1 = torch.sum(output1.float())
|
||||
sum_output2 = torch.sum(output2.float())
|
||||
logger.info(f"Rank {sp_rank} Sum of output1: {sum_output1.item()}")
|
||||
logger.info(f"Rank {sp_rank} Sum of output2: {sum_output2.item()}")
|
||||
|
||||
# Clean up
|
||||
logger.info("Cleaning up distributed environment")
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
logger.info("Test completed successfully")
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_hunyuanvideo_distributed()
|
||||
@@ -0,0 +1,201 @@
|
||||
from fastvideo.v1.models.encoders.llama import LlamaModel
|
||||
from fastvideo.v1.v0_reference_src.models.hunyuan.text_encoder import TextEncoder, load_text_encoder, load_tokenizer
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import argparse
|
||||
import numpy as np
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from transformers import AutoConfig
|
||||
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
def initialize_identical_weights(model1, model2, seed=42):
|
||||
"""Initialize both models with identical weights using a fixed seed for reproducibility."""
|
||||
# Get all parameters from both models
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
# Initialize each layer with identical values
|
||||
with torch.no_grad():
|
||||
# Initialize weights
|
||||
for name1, param1 in params1.items():
|
||||
if 'weight' in name1:
|
||||
# Set seed before each weight initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'weight' in name2:
|
||||
# Reset seed to get same initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
|
||||
# Initialize biases
|
||||
for name1, param1 in params1.items():
|
||||
if 'bias' in name1:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
param1.data = param1.data.to(torch.bfloat16)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'bias' in name2:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
param2.data = param2.data.to(torch.bfloat16)
|
||||
|
||||
logger.info("Both models initialized with identical weights in bfloat16")
|
||||
return model1, model2
|
||||
|
||||
def setup_args():
|
||||
parser = argparse.ArgumentParser(description='LLaMA Encoder Test')
|
||||
parser.add_argument('--model-path', type=str, default="meta-llama/Llama-2-7b-hf",
|
||||
help='Path to the LLaMA model')
|
||||
parser.add_argument('--precision', type=str, default="float16",
|
||||
help='Precision to use for the model (float32, float16, bfloat16)')
|
||||
return parser.parse_args()
|
||||
|
||||
def test_llama_encoder():
|
||||
init_distributed_environment(world_size=1, rank=0, distributed_init_method="env://", local_rank=0, backend="nccl")
|
||||
initialize_model_parallel(tensor_model_parallel_size=1, sequence_model_parallel_size=1, backend="nccl")
|
||||
args = setup_args()
|
||||
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
# Initialize the two model implementations
|
||||
logger.info(f"Loading models from {args.model_path}")
|
||||
model_path = "data/hunyuanvideo-community/HunyuanVideo/text_encoder"
|
||||
|
||||
hf_config = AutoConfig.from_pretrained(model_path)
|
||||
print(hf_config)
|
||||
|
||||
# Load our implementation using the loader from text_encoder/__init__.py
|
||||
model1, _ = load_text_encoder(
|
||||
text_encoder_type="llm",
|
||||
text_encoder_precision='fp16',
|
||||
text_encoder_path=model_path,
|
||||
logger=logger,
|
||||
device=device
|
||||
)
|
||||
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.v1.models.loader.loader import TextEncoderLoader
|
||||
loader = TextEncoderLoader()
|
||||
args.device_str = "cuda:0"
|
||||
model2 = loader.load_model(model_path, hf_config, args)
|
||||
|
||||
# Convert to float16 and move to device
|
||||
model2 = model2.to(torch.float16)
|
||||
model2 = model2.to(device)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights between the two models
|
||||
logger.info("Comparing model weights for sanity check...")
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
# Check number of parameters
|
||||
logger.info(f"Model1 has {len(params1)} parameters")
|
||||
logger.info(f"Model2 has {len(params2)} parameters")
|
||||
|
||||
# Compare a few key parameters
|
||||
weight_diffs = []
|
||||
# check if embed_tokens are the same
|
||||
print(model1.embed_tokens.weight.shape, model2.embed_tokens.weight.shape)
|
||||
assert torch.allclose(model1.embed_tokens.weight, model2.embed_tokens.weight)
|
||||
weights = ["layers.{}.input_layernorm.weight", "layers.{}.post_attention_layernorm.weight"]
|
||||
# for (name1, param1), (name2, param2) in zip(
|
||||
# sorted(params1.items()), sorted(params2.items())
|
||||
# ):
|
||||
for l in range(hf_config.num_hidden_layers):
|
||||
for w in weights:
|
||||
name1 = w.format(l)
|
||||
name2 = w.format(l)
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
print(type(p2))
|
||||
if "gate_up" in name2:
|
||||
print("skipping gate_up")
|
||||
continue
|
||||
try:
|
||||
logger.info(f"Parameter: {name1} vs {name2}")
|
||||
max_diff = torch.max(torch.abs(p1 - p2)).item()
|
||||
mean_diff = torch.mean(torch.abs(p1 - p2)).item()
|
||||
weight_diffs.append((name1, name2, max_diff, mean_diff))
|
||||
logger.info(f" Max diff: {max_diff}, Mean diff: {mean_diff}")
|
||||
except Exception as e:
|
||||
logger.info(f"Error comparing {name1} and {name2}: {e}")
|
||||
|
||||
tokenizer_path = "data/hunyuanvideo-community/HunyuanVideo/tokenizer"
|
||||
# Load tokenizer
|
||||
tokenizer, _ = load_tokenizer(
|
||||
tokenizer_type="llm",
|
||||
tokenizer_path=tokenizer_path,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
# Test with some sample prompts
|
||||
prompts = [
|
||||
"Once upon a time",
|
||||
# "The quick brown fox jumps over",
|
||||
# "In a galaxy far, far away"
|
||||
]
|
||||
|
||||
logger.info("Testing LLaMA encoder with sample prompts")
|
||||
|
||||
with torch.no_grad():
|
||||
for prompt in prompts:
|
||||
logger.info(f"Testing prompt: '{prompt}'")
|
||||
|
||||
# Tokenize the prompt
|
||||
tokens = tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=128,
|
||||
truncation=True,
|
||||
return_tensors="pt"
|
||||
).to(device)
|
||||
|
||||
# Get outputs from our implementation
|
||||
# filter out padding input_ids
|
||||
# tokens.input_ids = tokens.input_ids[tokens.attention_mask==1]
|
||||
# tokens.attention_mask = tokens.attention_mask[tokens.attention_mask==1]
|
||||
outputs1 = model1(
|
||||
input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
output_hidden_states=True
|
||||
)
|
||||
print("--------------------------------")
|
||||
logger.info(f"Testing model2")
|
||||
|
||||
# Get outputs from HuggingFace implementation
|
||||
outputs2 = model2(
|
||||
input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
output_hidden_states=True
|
||||
)
|
||||
|
||||
# Compare last hidden states
|
||||
last_hidden_state1 = outputs1.last_hidden_state[tokens.attention_mask==1]
|
||||
last_hidden_state2 = outputs2.last_hidden_state[tokens.attention_mask==1]
|
||||
|
||||
assert last_hidden_state1.shape == last_hidden_state2.shape, \
|
||||
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
|
||||
|
||||
max_diff_hidden = torch.max(torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
mean_diff_hidden = torch.mean(torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
|
||||
logger.info(f"Maximum difference in last hidden states: {max_diff_hidden.item()}")
|
||||
logger.info(f"Mean difference in last hidden states: {mean_diff_hidden.item()}")
|
||||
|
||||
|
||||
logger.info("Test passed! Both LLaMA encoder implementations produce similar outputs.")
|
||||
logger.info("Test completed successfully")
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_llama_encoder()
|
||||
@@ -0,0 +1,163 @@
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.distributed as dist
|
||||
import argparse
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size,
|
||||
destroy_model_parallel,
|
||||
destroy_distributed_environment,
|
||||
cleanup_dist_env_and_memory
|
||||
)
|
||||
from fastvideo.v1.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
RowParallelLinear
|
||||
)
|
||||
from fastvideo.v1.distributed.communication_op import (
|
||||
tensor_model_parallel_all_reduce,
|
||||
tensor_model_parallel_all_gather
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
class SimpleTPModel(nn.Module):
|
||||
"""A simple model that uses tensor parallelism."""
|
||||
|
||||
def __init__(self, hidden_size=1024, intermediate_size=4096):
|
||||
super().__init__()
|
||||
# Column parallel linear layer (splits output dimension)
|
||||
self.fc1 = ColumnParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_size=intermediate_size,
|
||||
bias=True,
|
||||
gather_output=False, # Don't gather output since we're passing to row parallel
|
||||
skip_bias_add=False
|
||||
)
|
||||
|
||||
# Row parallel linear layer (splits input dimension)
|
||||
self.fc2 = RowParallelLinear(
|
||||
input_size=intermediate_size,
|
||||
output_size=hidden_size,
|
||||
bias=True,
|
||||
input_is_parallel=True, # Input is already split from previous layer
|
||||
skip_bias_add=False
|
||||
)
|
||||
|
||||
self.activation = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
# Forward through column parallel layer
|
||||
hidden_states, _ = self.fc1(x)
|
||||
|
||||
# Apply activation
|
||||
hidden_states = self.activation(hidden_states)
|
||||
|
||||
# Forward through row parallel layer
|
||||
output, _ = self.fc2(hidden_states)
|
||||
|
||||
return output
|
||||
|
||||
def initialize_random_weights(model, seed=42):
|
||||
"""Initialize the model with random weights using a fixed seed for reproducibility."""
|
||||
# Set seed for reproducibility
|
||||
torch.manual_seed(seed)
|
||||
|
||||
# Initialize weights for each layer
|
||||
with torch.no_grad():
|
||||
# For ColumnParallelLinear layers
|
||||
if hasattr(model, 'fc1'):
|
||||
nn.init.normal_(model.fc1.weight, mean=0.0, std=0.02)
|
||||
if model.fc1.bias is not None:
|
||||
nn.init.zeros_(model.fc1.bias)
|
||||
|
||||
# For RowParallelLinear layers
|
||||
if hasattr(model, 'fc2'):
|
||||
nn.init.normal_(model.fc2.weight, mean=0.0, std=0.02)
|
||||
if model.fc2.bias is not None:
|
||||
nn.init.zeros_(model.fc2.bias)
|
||||
|
||||
logger.info("Model initialized with random weights")
|
||||
return model
|
||||
|
||||
def setup_args():
|
||||
parser = argparse.ArgumentParser(description='Simple Tensor Parallelism Example')
|
||||
parser.add_argument('--tensor-model-parallel-size', type=int, default=8,
|
||||
help='Degree of tensor model parallelism')
|
||||
parser.add_argument('--batch-size', type=int, default=8,
|
||||
help='Batch size for the example')
|
||||
parser.add_argument('--hidden-size', type=int, default=1024,
|
||||
help='Hidden size for the model')
|
||||
parser.add_argument('--intermediate-size', type=int, default=4096,
|
||||
help='Intermediate size for the model')
|
||||
return parser.parse_args()
|
||||
|
||||
def main():
|
||||
args = setup_args()
|
||||
|
||||
# Initialize distributed environment
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
|
||||
logger.info(f"Initializing process: rank={rank}, local_rank={local_rank}, world_size={world_size}")
|
||||
|
||||
init_distributed_environment(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank
|
||||
)
|
||||
|
||||
# Initialize tensor model parallel groups
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=args.tensor_model_parallel_size
|
||||
)
|
||||
|
||||
# Get tensor parallel info
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
tp_world_size = get_tensor_model_parallel_world_size()
|
||||
|
||||
logger.info(f"Process rank {rank} initialized with TP rank {tp_rank} in TP world size {tp_world_size}")
|
||||
|
||||
# Create a simple model
|
||||
model = SimpleTPModel(
|
||||
hidden_size=args.hidden_size,
|
||||
intermediate_size=args.intermediate_size
|
||||
)
|
||||
|
||||
# Initialize with random weights
|
||||
model = initialize_random_weights(model)
|
||||
|
||||
# Create a random input tensor
|
||||
batch_size = args.batch_size
|
||||
hidden_size = args.hidden_size
|
||||
x = torch.randn(batch_size, hidden_size, dtype=torch.float)
|
||||
|
||||
# Move to GPU if available
|
||||
device = torch.device(f"cuda:{local_rank}" if torch.cuda.is_available() else "cpu")
|
||||
model = model.to(device)
|
||||
x = x.to(device)
|
||||
|
||||
# Forward pass
|
||||
logger.info(f"Running forward pass on TP rank {tp_rank}")
|
||||
with torch.no_grad():
|
||||
output = model(x)
|
||||
|
||||
# Print output shape and statistics
|
||||
logger.info(f"Output shape: {output.shape}")
|
||||
logger.info(f"Output mean: {output.mean().item()}, std: {output.std().item()}")
|
||||
|
||||
# Clean up
|
||||
logger.info("Cleaning up distributed environment")
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
logger.info("Example completed successfully")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,312 @@
|
||||
import torch
|
||||
import fastvideo.v1.envs as envs
|
||||
import inspect
|
||||
from fastvideo.v1.logger import init_logger
|
||||
import argparse
|
||||
import math
|
||||
import sys
|
||||
from typing import List, Dict, Union, Type, Any, TypeVar
|
||||
import yaml
|
||||
from functools import wraps
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
# TODO(will): used to convert inference_args.precision to torch.dtype. Find a
|
||||
# cleaner way to do this.
|
||||
PRECISION_TO_TYPE = {
|
||||
"fp32": torch.float32,
|
||||
"fp16": torch.float16,
|
||||
"bf16": torch.bfloat16,
|
||||
}
|
||||
|
||||
|
||||
def find_nccl_library() -> str:
|
||||
"""
|
||||
We either use the library file specified by the `VLLM_NCCL_SO_PATH`
|
||||
environment variable, or we find the library file brought by PyTorch.
|
||||
After importing `torch`, `libnccl.so.2` or `librccl.so.1` can be
|
||||
found by `ctypes` automatically.
|
||||
"""
|
||||
so_file = envs.FASTVIDEO_NCCL_SO_PATH
|
||||
|
||||
# manually load the nccl library
|
||||
if so_file:
|
||||
logger.info(
|
||||
"Found nccl from environment variable FASTVIDEO_NCCL_SO_PATH=%s",
|
||||
so_file)
|
||||
else:
|
||||
if torch.version.cuda is not None:
|
||||
so_file = "libnccl.so.2"
|
||||
elif torch.version.hip is not None:
|
||||
so_file = "librccl.so.1"
|
||||
else:
|
||||
raise ValueError("NCCL only supports CUDA and ROCm backends.")
|
||||
logger.info("Found nccl from library %s", so_file)
|
||||
return so_file
|
||||
|
||||
prev_set_stream = torch.cuda.set_stream
|
||||
|
||||
_current_stream = None
|
||||
|
||||
|
||||
def _patched_set_stream(stream: torch.cuda.Stream) -> None:
|
||||
global _current_stream
|
||||
_current_stream = stream
|
||||
prev_set_stream(stream)
|
||||
|
||||
|
||||
torch.cuda.set_stream = _patched_set_stream
|
||||
|
||||
|
||||
def current_stream() -> torch.cuda.Stream:
|
||||
"""
|
||||
replace `torch.cuda.current_stream()` with `vllm.utils.current_stream()`.
|
||||
it turns out that `torch.cuda.current_stream()` is quite expensive,
|
||||
as it will construct a new stream object at each call.
|
||||
here we patch `torch.cuda.set_stream` to keep track of the current stream
|
||||
directly, so that we can avoid calling `torch.cuda.current_stream()`.
|
||||
|
||||
the underlying hypothesis is that we do not call `torch._C._cuda_setStream`
|
||||
from C/C++ code.
|
||||
"""
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
global _current_stream
|
||||
if _current_stream is None:
|
||||
# when this function is called before any stream is set,
|
||||
# we return the default stream.
|
||||
# On ROCm using the default 0 stream in combination with RCCL
|
||||
# is hurting performance. Therefore creating a dedicated stream
|
||||
# per process
|
||||
_current_stream = torch.cuda.Stream() if current_platform.is_rocm(
|
||||
) else torch.cuda.current_stream()
|
||||
return _current_stream
|
||||
|
||||
|
||||
class StoreBoolean(argparse.Action):
|
||||
|
||||
def __call__(self, parser, namespace, values, option_string=None):
|
||||
if values.lower() == "true":
|
||||
setattr(namespace, self.dest, True)
|
||||
elif values.lower() == "false":
|
||||
setattr(namespace, self.dest, False)
|
||||
else:
|
||||
raise ValueError(f"Invalid boolean value: {values}. "
|
||||
"Expected 'true' or 'false'.")
|
||||
|
||||
|
||||
class SortedHelpFormatter(argparse.HelpFormatter):
|
||||
"""SortedHelpFormatter that sorts arguments by their option strings."""
|
||||
|
||||
def add_arguments(self, actions):
|
||||
actions = sorted(actions, key=lambda x: x.option_strings)
|
||||
super().add_arguments(actions)
|
||||
|
||||
|
||||
class FlexibleArgumentParser(argparse.ArgumentParser):
|
||||
"""ArgumentParser that allows both underscore and dash in names."""
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
# Set the default 'formatter_class' to SortedHelpFormatter
|
||||
if 'formatter_class' not in kwargs:
|
||||
kwargs['formatter_class'] = SortedHelpFormatter
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def parse_args(self, args=None, namespace=None):
|
||||
if args is None:
|
||||
args = sys.argv[1:]
|
||||
|
||||
if '--config' in args:
|
||||
args = self._pull_args_from_config(args)
|
||||
|
||||
# Convert underscores to dashes and vice versa in argument names
|
||||
processed_args = []
|
||||
for arg in args:
|
||||
if arg.startswith('--'):
|
||||
if '=' in arg:
|
||||
key, value = arg.split('=', 1)
|
||||
key = '--' + key[len('--'):].replace('_', '-')
|
||||
processed_args.append(f'{key}={value}')
|
||||
else:
|
||||
processed_args.append('--' +
|
||||
arg[len('--'):].replace('_', '-'))
|
||||
elif arg.startswith('-O') and arg != '-O' and len(arg) == 2:
|
||||
# allow -O flag to be used without space, e.g. -O3
|
||||
processed_args.append('-O')
|
||||
processed_args.append(arg[2:])
|
||||
else:
|
||||
processed_args.append(arg)
|
||||
|
||||
return super().parse_args(processed_args, namespace)
|
||||
|
||||
def _pull_args_from_config(self, args: List[str]) -> List[str]:
|
||||
"""Method to pull arguments specified in the config file
|
||||
into the command-line args variable.
|
||||
|
||||
The arguments in config file will be inserted between
|
||||
the argument list.
|
||||
|
||||
example:
|
||||
```yaml
|
||||
port: 12323
|
||||
tensor-parallel-size: 4
|
||||
```
|
||||
```python
|
||||
$: vllm {serve,chat,complete} "facebook/opt-12B" \
|
||||
--config config.yaml -tp 2
|
||||
$: args = [
|
||||
"serve,chat,complete",
|
||||
"facebook/opt-12B",
|
||||
'--config', 'config.yaml',
|
||||
'-tp', '2'
|
||||
]
|
||||
$: args = [
|
||||
"serve,chat,complete",
|
||||
"facebook/opt-12B",
|
||||
'--port', '12323',
|
||||
'--tensor-parallel-size', '4',
|
||||
'-tp', '2'
|
||||
]
|
||||
```
|
||||
|
||||
Please note how the config args are inserted after the sub command.
|
||||
this way the order of priorities is maintained when these are args
|
||||
parsed by super().
|
||||
"""
|
||||
assert args.count(
|
||||
'--config') <= 1, "More than one config file specified!"
|
||||
|
||||
index = args.index('--config')
|
||||
if index == len(args) - 1:
|
||||
raise ValueError("No config file specified! \
|
||||
Please check your command-line arguments.")
|
||||
|
||||
file_path = args[index + 1]
|
||||
|
||||
config_args = self._load_config_file(file_path)
|
||||
|
||||
# 0th index is for {serve,chat,complete}
|
||||
# followed by model_tag (only for serve)
|
||||
# followed by config args
|
||||
# followed by rest of cli args.
|
||||
# maintaining this order will enforce the precedence
|
||||
# of cli > config > defaults
|
||||
if args[0] == "serve":
|
||||
if index == 1:
|
||||
raise ValueError(
|
||||
"No model_tag specified! Please check your command-line"
|
||||
" arguments.")
|
||||
args = [args[0]] + [
|
||||
args[1]
|
||||
] + config_args + args[2:index] + args[index + 2:]
|
||||
else:
|
||||
args = [args[0]] + config_args + args[1:index] + args[index + 2:]
|
||||
|
||||
return args
|
||||
|
||||
def _load_config_file(self, file_path: str) -> List[str]:
|
||||
"""Loads a yaml file and returns the key value pairs as a
|
||||
flattened list with argparse like pattern
|
||||
```yaml
|
||||
port: 12323
|
||||
tensor-parallel-size: 4
|
||||
```
|
||||
returns:
|
||||
processed_args: list[str] = [
|
||||
'--port': '12323',
|
||||
'--tensor-parallel-size': '4'
|
||||
]
|
||||
|
||||
"""
|
||||
|
||||
extension: str = file_path.split('.')[-1]
|
||||
if extension not in ('yaml', 'yml'):
|
||||
raise ValueError(
|
||||
"Config file must be of a yaml/yml type.\
|
||||
%s supplied", extension)
|
||||
|
||||
# only expecting a flat dictionary of atomic types
|
||||
processed_args: List[str] = []
|
||||
|
||||
config: Dict[str, Union[int, str]] = {}
|
||||
try:
|
||||
with open(file_path) as config_file:
|
||||
config = yaml.safe_load(config_file)
|
||||
except Exception as ex:
|
||||
logger.error(
|
||||
"Unable to read the config file at %s. \
|
||||
Make sure path is correct", file_path)
|
||||
raise ex
|
||||
|
||||
store_boolean_arguments = [
|
||||
action.dest for action in self._actions
|
||||
if isinstance(action, StoreBoolean)
|
||||
]
|
||||
|
||||
for key, value in config.items():
|
||||
if isinstance(value, bool) and key not in store_boolean_arguments:
|
||||
if value:
|
||||
processed_args.append('--' + key)
|
||||
else:
|
||||
processed_args.append('--' + key)
|
||||
processed_args.append(str(value))
|
||||
|
||||
return processed_args
|
||||
|
||||
|
||||
def warn_for_unimplemented_methods(cls: Type[T]) -> Type[T]:
|
||||
"""
|
||||
A replacement for `abc.ABC`.
|
||||
When we use `abc.ABC`, subclasses will fail to instantiate
|
||||
if they do not implement all abstract methods.
|
||||
Here, we only require `raise NotImplementedError` in the
|
||||
base class, and log a warning if the method is not implemented
|
||||
in the subclass.
|
||||
"""
|
||||
|
||||
original_init = cls.__init__
|
||||
|
||||
def find_unimplemented_methods(self: object):
|
||||
unimplemented_methods = []
|
||||
for attr_name in dir(self):
|
||||
# bypass inner method
|
||||
if attr_name.startswith('_'):
|
||||
continue
|
||||
|
||||
try:
|
||||
attr = getattr(self, attr_name)
|
||||
# get the func of callable method
|
||||
if callable(attr):
|
||||
attr_func = attr.__func__
|
||||
except AttributeError:
|
||||
continue
|
||||
src = inspect.getsource(attr_func)
|
||||
if "NotImplementedError" in src:
|
||||
unimplemented_methods.append(attr_name)
|
||||
if unimplemented_methods:
|
||||
method_names = ','.join(unimplemented_methods)
|
||||
msg = (f"Methods {method_names} not implemented in {self}")
|
||||
logger.warning(msg)
|
||||
|
||||
@wraps(original_init)
|
||||
def wrapped_init(self, *args, **kwargs) -> None:
|
||||
original_init(self, *args, **kwargs)
|
||||
find_unimplemented_methods(self)
|
||||
|
||||
type.__setattr__(cls, '__init__', wrapped_init)
|
||||
return cls
|
||||
|
||||
|
||||
def align_to(value, alignment):
|
||||
"""align height, width according to alignment
|
||||
|
||||
Args:
|
||||
value (int): height or width
|
||||
alignment (int): target alignment factor
|
||||
|
||||
Returns:
|
||||
int: the aligned value
|
||||
"""
|
||||
return int(math.ceil(value / alignment) * alignment)
|
||||
+1
-1
@@ -12,7 +12,7 @@ import torch
|
||||
import torchvision
|
||||
from cog import BasePredictor, Input, Path
|
||||
from einops import rearrange
|
||||
|
||||
from transformers import LlamaForCausalLM
|
||||
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
|
||||
|
||||
MODEL_CACHE = 'FastHunyuan'
|
||||
|
||||
+1
-1
@@ -76,4 +76,4 @@ skip ="./data,./wandb,./csrc/sliding_tile_attention/tk"
|
||||
column_limit = 120
|
||||
|
||||
[tool.isort]
|
||||
line_length = 120
|
||||
line_length = 120
|
||||
@@ -0,0 +1 @@
|
||||
huggingface_hub
|
||||
@@ -0,0 +1,28 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=8
|
||||
export TOKENIZERS_PARALLELISM=true
|
||||
export MODEL_BASE=data/hunyuan_diffusers
|
||||
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/v1/sample/v1_fastvideo_inference.py \
|
||||
--use-v1-transformer \
|
||||
--use-v1-vae \
|
||||
--use-v1-text-encoder \
|
||||
--sp_size $num_gpus \
|
||||
--tp_size $num_gpus \
|
||||
--height 768 \
|
||||
--width 1280 \
|
||||
--num_frames 125 \
|
||||
--num_inference_steps 50 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow_shift 7 \
|
||||
--prompt_path ./assets/prompt.txt \
|
||||
--seed 12345 \
|
||||
--output_path outputs_video/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--vae-sp
|
||||
@@ -1,252 +0,0 @@
|
||||
# Copyright (c) Tile-AI Corporation.
|
||||
# Licensed under the MIT License.
|
||||
import math
|
||||
import torch
|
||||
|
||||
import tilelang
|
||||
import tilelang.language as T
|
||||
import torch.nn.functional as F
|
||||
|
||||
def get_sta_mask(x, canvas_size=(32, 48, 80), tile_size=(4, 8, 8), kernel_size=(1, 1, 1), has_text=False):
|
||||
bsz, num_head, downsample_len, _ = x.shape
|
||||
device = x.device
|
||||
CT, CH, CW = canvas_size
|
||||
TT, TH, TW = tile_size
|
||||
NT, NH, NW = CT // TT, CH // TH, CW // TW
|
||||
KT, KH, KW = kernel_size
|
||||
DT, DH, DW = KT // 2, KH // 2, KW // 2
|
||||
|
||||
dense_mask = torch.full([bsz, num_head, downsample_len, downsample_len],
|
||||
False,
|
||||
dtype=torch.bool,
|
||||
device=device)
|
||||
indices = torch.arange(downsample_len, device=device)
|
||||
q_t = (indices // (NT * NH * NW)).unsqueeze(-1)
|
||||
q_h = ((indices // (NW)) % NH).unsqueeze(-1)
|
||||
q_w = ((indices // 1) % NW).unsqueeze(-1)
|
||||
|
||||
q_t = torch.clamp(q_t, DT, NT-DT-1)
|
||||
q_h = torch.clamp(q_h, DH, NH-DH-1)
|
||||
q_w = torch.clamp(q_w, DW, NW-DW-1)
|
||||
|
||||
k_indices = torch.arange(downsample_len, device=device)
|
||||
k_t = (k_indices // (NT * NH * NW)).unsqueeze(0)
|
||||
k_h = ((k_indices // (NW)) % NH).unsqueeze(0)
|
||||
k_w = ((k_indices // 1) % NW).unsqueeze(0)
|
||||
|
||||
t_dist = torch.abs(q_t - k_t)
|
||||
h_dist = torch.abs(q_h - k_h)
|
||||
w_dist = torch.abs(q_w - k_w)
|
||||
mask = (t_dist <= DT) & (h_dist <= DH) & (w_dist <= DW)
|
||||
|
||||
for b in range(bsz):
|
||||
for h in range(num_head):
|
||||
dense_mask[b, h] = mask
|
||||
|
||||
# for text mask
|
||||
if has_text:
|
||||
text_start = downsample_len - 3
|
||||
dense_mask[:, :, text_start:, :text_start] = True
|
||||
dense_mask[:, :, text_start:, text_start:] = torch.tril(
|
||||
torch.ones(3, 3, device=device, dtype=torch.bool)
|
||||
)
|
||||
return dense_mask
|
||||
|
||||
def blocksparse_flashattn(batch, heads, seq_len, dim, downsample_len, is_causal):
|
||||
block_M = 64
|
||||
block_N = 64
|
||||
num_stages = 1
|
||||
threads = 128
|
||||
scale = (1.0 / dim)**0.5 * 1.44269504 # log2(e)
|
||||
shape = [batch, heads, seq_len, dim]
|
||||
block_mask_shape = [batch, heads, downsample_len, downsample_len]
|
||||
|
||||
dtype = "float16"
|
||||
accum_dtype = "float"
|
||||
block_mask_dtype = "bool"
|
||||
|
||||
def kernel_func(block_M, block_N, num_stages, threads):
|
||||
|
||||
@T.macro
|
||||
def MMA0(
|
||||
K: T.Buffer(shape, dtype),
|
||||
Q_shared: T.Buffer([block_M, dim], dtype),
|
||||
K_shared: T.Buffer([block_N, dim], dtype),
|
||||
acc_s: T.Buffer([block_M, block_N], accum_dtype),
|
||||
k: T.int32,
|
||||
bx: T.int32,
|
||||
by: T.int32,
|
||||
bz: T.int32,
|
||||
):
|
||||
T.copy(K[bz, by, k * block_N:(k + 1) * block_N, :], K_shared)
|
||||
if is_causal:
|
||||
for i, j in T.Parallel(block_M, block_N):
|
||||
acc_s[i, j] = T.if_then_else(bx * block_M + i >= k * block_N + j, 0,
|
||||
-T.infinity(acc_s.dtype))
|
||||
else:
|
||||
T.clear(acc_s)
|
||||
T.gemm(Q_shared, K_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullRow)
|
||||
|
||||
@T.macro
|
||||
def MMA1(
|
||||
V: T.Buffer(shape, dtype),
|
||||
V_shared: T.Buffer([block_M, dim], dtype),
|
||||
acc_s_cast: T.Buffer([block_M, block_N], dtype),
|
||||
acc_o: T.Buffer([block_M, dim], accum_dtype),
|
||||
k: T.int32,
|
||||
by: T.int32,
|
||||
bz: T.int32,
|
||||
):
|
||||
T.copy(V[bz, by, k * block_N:(k + 1) * block_N, :], V_shared)
|
||||
T.gemm(acc_s_cast, V_shared, acc_o, policy=T.GemmWarpPolicy.FullRow)
|
||||
|
||||
@T.macro
|
||||
def Softmax(
|
||||
acc_s: T.Buffer([block_M, block_N], accum_dtype),
|
||||
acc_s_cast: T.Buffer([block_M, block_N], dtype),
|
||||
scores_max: T.Buffer([block_M], accum_dtype),
|
||||
scores_max_prev: T.Buffer([block_M], accum_dtype),
|
||||
scores_scale: T.Buffer([block_M], accum_dtype),
|
||||
scores_sum: T.Buffer([block_M], accum_dtype),
|
||||
logsum: T.Buffer([block_M], accum_dtype),
|
||||
):
|
||||
T.copy(scores_max, scores_max_prev)
|
||||
T.fill(scores_max, -T.infinity(accum_dtype))
|
||||
T.reduce_max(acc_s, scores_max, dim=1, clear=False)
|
||||
# To do causal softmax, we need to set the scores_max to 0 if it is -inf
|
||||
# This process is called Check_inf in FlashAttention3 code, and it only need to be done
|
||||
# in the first ceil_div(kBlockM, kBlockN) steps.
|
||||
# for i in T.Parallel(block_M):
|
||||
# scores_max[i] = T.if_then_else(scores_max[i] == -T.infinity(accum_dtype), 0, scores_max[i])
|
||||
for i in T.Parallel(block_M):
|
||||
scores_scale[i] = T.exp2(scores_max_prev[i] * scale - scores_max[i] * scale)
|
||||
for i, j in T.Parallel(block_M, block_N):
|
||||
# Instead of computing exp(x - max), we compute exp2(x * log_2(e) -
|
||||
# max * log_2(e)) This allows the compiler to use the ffma
|
||||
# instruction instead of fadd and fmul separately.
|
||||
acc_s[i, j] = T.exp2(acc_s[i, j] * scale - scores_max[i] * scale)
|
||||
T.reduce_sum(acc_s, scores_sum, dim=1)
|
||||
for i in T.Parallel(block_M):
|
||||
logsum[i] = logsum[i] * scores_scale[i] + scores_sum[i]
|
||||
T.copy(acc_s, acc_s_cast)
|
||||
|
||||
@T.macro
|
||||
def Rescale(
|
||||
acc_o: T.Buffer([block_M, dim], accum_dtype),
|
||||
scores_scale: T.Buffer([block_M], accum_dtype),
|
||||
):
|
||||
for i, j in T.Parallel(block_M, dim):
|
||||
acc_o[i, j] *= scores_scale[i]
|
||||
|
||||
@T.prim_func
|
||||
def main(
|
||||
Q: T.Buffer(shape, dtype),
|
||||
K: T.Buffer(shape, dtype),
|
||||
V: T.Buffer(shape, dtype),
|
||||
BlockSparseMask: T.Buffer(block_mask_shape, block_mask_dtype),
|
||||
Output: T.Buffer(shape, dtype),
|
||||
):
|
||||
with T.Kernel(
|
||||
T.ceildiv(seq_len, block_M), heads, batch, threads=threads) as (bx, by, bz):
|
||||
Q_shared = T.alloc_shared([block_M, dim], dtype)
|
||||
K_shared = T.alloc_shared([block_N, dim], dtype)
|
||||
V_shared = T.alloc_shared([block_N, dim], dtype)
|
||||
O_shared = T.alloc_shared([block_M, dim], dtype)
|
||||
acc_s = T.alloc_fragment([block_M, block_N], accum_dtype)
|
||||
acc_s_cast = T.alloc_fragment([block_M, block_N], dtype)
|
||||
acc_o = T.alloc_fragment([block_M, dim], accum_dtype)
|
||||
scores_max = T.alloc_fragment([block_M], accum_dtype)
|
||||
scores_max_prev = T.alloc_fragment([block_M], accum_dtype)
|
||||
scores_scale = T.alloc_fragment([block_M], accum_dtype)
|
||||
scores_sum = T.alloc_fragment([block_M], accum_dtype)
|
||||
logsum = T.alloc_fragment([block_M], accum_dtype)
|
||||
block_mask = T.alloc_local([downsample_len], block_mask_dtype)
|
||||
|
||||
T.copy(Q[bz, by, bx * block_M:(bx + 1) * block_M, :], Q_shared)
|
||||
T.fill(acc_o, 0)
|
||||
T.fill(logsum, 0)
|
||||
T.fill(scores_max, -T.infinity(accum_dtype))
|
||||
|
||||
for vj in T.serial(downsample_len):
|
||||
block_mask[vj] = BlockSparseMask[bz, by, bx, vj]
|
||||
|
||||
loop_range = (
|
||||
T.min(T.ceildiv(seq_len, block_N), T.ceildiv(
|
||||
(bx + 1) * block_M, block_N)) if is_causal else T.ceildiv(seq_len, block_N))
|
||||
|
||||
for k in T.Pipelined(loop_range, num_stages=num_stages):
|
||||
if block_mask[k] != 0:
|
||||
MMA0(K, Q_shared, K_shared, acc_s, k, bx, by, bz)
|
||||
Softmax(acc_s, acc_s_cast, scores_max, scores_max_prev, scores_scale,
|
||||
scores_sum, logsum)
|
||||
Rescale(acc_o, scores_scale)
|
||||
MMA1(V, V_shared, acc_s_cast, acc_o, k, by, bz)
|
||||
for i, j in T.Parallel(block_M, dim):
|
||||
acc_o[i, j] /= logsum[i]
|
||||
T.copy(acc_o, O_shared)
|
||||
T.copy(O_shared, Output[bz, by, bx * block_M:(bx + 1) * block_M, :])
|
||||
|
||||
return main
|
||||
|
||||
return kernel_func(block_M, block_N, num_stages, threads)
|
||||
|
||||
|
||||
def test_sta_attention():
|
||||
# Config
|
||||
BATCH, N_HEADS, SEQ_LEN, D_HEAD = 1, 24, 2048, 128
|
||||
torch.manual_seed(0)
|
||||
|
||||
# Create inputs
|
||||
q = torch.randn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, device='cuda', dtype=torch.float16)
|
||||
k = torch.randn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, device='cuda', dtype=torch.float16)
|
||||
v = torch.randn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, device='cuda', dtype=torch.float16)
|
||||
|
||||
sm_scale = 1.0 / (D_HEAD**0.5)
|
||||
|
||||
# Create sparse mask (downsampled to block level)
|
||||
canvas_size = (32, 48, 80)
|
||||
tile_size = (4, 8, 8)
|
||||
delta_size = (2, 2, 2)
|
||||
BLOCK = tile_size[0] * tile_size[1] * tile_size[2]
|
||||
downsample_factor = BLOCK
|
||||
downsample_len = math.ceil(SEQ_LEN / downsample_factor)
|
||||
x_ds = torch.randn([BATCH, N_HEADS, downsample_len, downsample_len],
|
||||
device='cuda',
|
||||
dtype=torch.bfloat16)
|
||||
block_mask = get_sta_mask(x_ds, canvas_size, tile_size, delta_size, has_text=True)
|
||||
# print mask density
|
||||
print("mask density", block_mask.sum() / block_mask.numel())
|
||||
print("block_mask", block_mask)
|
||||
|
||||
# Run Triton kernel
|
||||
program = blocksparse_flashattn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, downsample_len, is_causal=True)
|
||||
kernel = tilelang.compile(program, out_idx=[4])
|
||||
|
||||
cuda_source = kernel.get_kernel_source()
|
||||
print("Generated CUDA kernel:\n", cuda_source)
|
||||
|
||||
tilelang_output = kernel(q, k, v, block_mask)
|
||||
|
||||
if True:
|
||||
# Compute reference
|
||||
# Expand block mask to full attention matrix
|
||||
full_mask = torch.kron(block_mask.float(), torch.ones(BLOCK, BLOCK, device='cuda'))
|
||||
full_mask = full_mask[..., :SEQ_LEN, :SEQ_LEN].bool()
|
||||
full_mask = full_mask & torch.tril(torch.ones_like(full_mask)) # Apply causal
|
||||
|
||||
# PyTorch reference implementation
|
||||
attn = torch.einsum('bhsd,bhtd->bhst', q, k) * sm_scale
|
||||
attn = attn.masked_fill(~full_mask, float('-inf'))
|
||||
attn = F.softmax(attn, dim=-1)
|
||||
ref_output = torch.einsum('bhst,bhtd->bhsd', attn, v)
|
||||
|
||||
print("ref_output", ref_output)
|
||||
print("tilelang_output", tilelang_output)
|
||||
|
||||
# Verify accuracy
|
||||
torch.testing.assert_close(tilelang_output, ref_output, atol=1e-2, rtol=1e-2)
|
||||
print("Pass topk sparse attention test with qlen == klen")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_sta_attention()
|
||||
Reference in New Issue
Block a user