Compare commits

..
Author SHA1 Message Date
Zihang-He 702e77c937 edit to attention 2025-05-23 01:39:30 +00:00
Zihang-He 2d43f78d04 compared tokenizer 2025-05-20 03:09:40 +00:00
Zihang-He e21a124f06 added bert encoder 2025-05-19 03:23:31 +00:00
24 changed files with 689 additions and 91 deletions
+1 -1
View File
@@ -8,7 +8,7 @@ body:
attributes:
label: Environment
description: |
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
Please share your environment with us. You can run the command **python fastvideo/utils/env_utils.py** and copy-paste its output below.
placeholder: FastVideo version, platform, python version, cuda version...
validations:
required: true
+2 -4
View File
@@ -77,8 +77,6 @@ jobs:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/transformers/**'
- 'fastvideo/v1/layers/**'
- 'fastvideo/v1/attention/**'
encoder-test:
needs: change-filter
@@ -143,8 +141,8 @@ jobs:
fail-fast: false
matrix:
python-version: [
# {version: "3.10", tag: "latest"},
# {version: "3.11", tag: "py3.11-latest"},
{version: "3.10", tag: "latest"},
{version: "3.11", tag: "py3.11-latest"},
{version: "3.12", tag: "py3.12-latest"}
]
uses: ./.github/workflows/runpod-test.yml
@@ -70,8 +70,6 @@ DEFAULT_CONDA_PATTERNS = {
"optree",
"nccl",
"transformers",
"accelerate",
"peft",
"zmq",
"nvidia",
"pynvml",
@@ -87,8 +85,6 @@ DEFAULT_PIP_PATTERNS = {
"onnx",
"nccl",
"transformers",
"accelerate",
"peft",
"zmq",
"nvidia",
"pynvml",
+2 -1
View File
@@ -2,6 +2,7 @@ import torch
from flex_sta_ref import get_sliding_tile_attention_mask
from st_attn import sliding_tile_attention
from torch.nn.attention.flex_attention import flex_attention
# from flash_attn_interface import flash_attn_func
from tqdm import tqdm
flex_attention = torch.compile(flex_attention, dynamic=False)
@@ -22,7 +23,7 @@ def h100_fwd_kernel_test(Q, K, V, kernel_size):
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.linalg.norm(tensor, dim=-1, keepdim=True)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
+1 -2
View File
@@ -1,6 +1,5 @@
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
from fastvideo.version import __version__
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam"]
+2 -2
View File
@@ -237,7 +237,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
type=str,
default="540p",
choices=["540p", "720p"],
help="The resolution of the model.",
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--load-key",
@@ -361,7 +361,7 @@ def add_parallel_args(parser: argparse.ArgumentParser):
"--ring-degree",
type=int,
default=1,
help="Ring degree.",
help="Ulysses degree.",
)
return parser
+1 -1
View File
@@ -17,7 +17,7 @@ from fastvideo.models.hunyuan.vae import load_vae
from fastvideo.utils.parallel_states import nccl_info
class Inference:
class Inference(object):
def __init__(
self,
+1 -1
View File
@@ -41,7 +41,7 @@ def get_rewrite_prompt(ori_prompt, mode="Normal"):
elif mode == "Master":
prompt = master_mode_prompt.format(input=ori_prompt)
else:
raise Exception("Only supports Normal and Master mode, but got {}".format(mode))
raise Exception("Only supports Normal and Normal", mode)
return prompt
@@ -16,7 +16,7 @@ class HunyuanClip(nn.Module):
self.max_length = max_length
self.tokenizer = BertTokenizer.from_pretrained(os.path.join(model_dir, 'tokenizer'))
self.text_encoder = BertModel.from_pretrained(os.path.join(model_dir, 'clip_text_encoder'))
self.text_encoder = BertModel.from_pretrained(os.path.join(model_dir, 'text_encoder'))
@torch.no_grad
def forward(self, prompts, with_mask=True):
@@ -31,6 +31,6 @@ class HunyuanClip(nn.Module):
)
prompt_embeds = self.text_encoder(
text_inputs.input_ids.to(self.device),
attention_mask=text_inputs.attention_mask.to(self.device) if with_mask else None,
# attention_mask=text_inputs.attention_mask.to(self.device) if with_mask else None,
)
return prompt_embeds.last_hidden_state, prompt_embeds.pooler_output
return prompt_embeds.last_hidden_state, text_inputs
@@ -267,25 +267,25 @@ class Step1Model(PreTrainedModel):
class STEP1TextEncoder(torch.nn.Module):
def __init__(self, model_dir, max_length=320):
super()
super(STEP1TextEncoder, self).__init__()
self.max_length = max_length
self.text_tokenizer = Wrapped_StepChatTokenizer(os.path.join(model_dir, 'step1_chat_tokenizer.model'))
text_encoder = Step1Model.from_pretrained(model_dir)
self.text_encoder = text_encoder.eval().to(torch.bfloat16)
@torch.no_grad
@torch.autocast(device_type='cuda', dtype=torch.bfloat16)
def forward(self, prompts, with_mask=True, max_length=None):
self.device = next(self.text_encoder.parameters()).device
if type(prompts) is str:
prompts = [prompts]
with torch.no_grad(), torch.cuda.amp.autocast(dtype=torch.bfloat16):
if type(prompts) is str:
prompts = [prompts]
txt_tokens = self.text_tokenizer(prompts,
max_length=max_length or self.max_length,
padding="max_length",
truncation=True,
return_tensors="pt")
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
txt_tokens = self.text_tokenizer(prompts,
max_length=max_length or self.max_length,
padding="max_length",
truncation=True,
return_tensors="pt")
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
attention_mask=txt_tokens.attention_mask.to(self.device) if with_mask else None)
y_mask = txt_tokens.attention_mask
y_mask = txt_tokens.attention_mask
return y.transpose(0, 1), y_mask
+38
View File
@@ -0,0 +1,38 @@
import platform
import accelerate
import peft
import torch
import transformers
from transformers.utils import is_torch_cuda_available, is_torch_npu_available
VERSION = "1.2.0"
if __name__ == "__main__":
info = {
"FastVideo version": VERSION,
"Platform": platform.platform(),
"Python version": platform.python_version(),
"PyTorch version": torch.__version__,
"Transformers version": transformers.__version__,
"Accelerate version": accelerate.__version__,
"PEFT version": peft.__version__,
}
if is_torch_cuda_available():
info["PyTorch version"] += " (GPU)"
info["GPU type"] = torch.cuda.get_device_name()
if is_torch_npu_available():
info["PyTorch version"] += " (NPU)"
info["NPU type"] = torch.npu.get_device_name()
info["CANN version"] = torch.version.cann # codespell:ignore
try:
import bitsandbytes
info["Bitsandbytes version"] = bitsandbytes.__version__
except Exception:
pass
print("\n" + "\n".join([f"- {key}: {value}" for key, value in info.items()]) + "\n")
+3
View File
@@ -84,6 +84,9 @@ class DistributedAttention(nn.Module):
# 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()
@@ -6,9 +6,10 @@ from fastvideo.v1.configs.models.encoders.clip import (CLIPTextConfig,
CLIPVisionConfig)
from fastvideo.v1.configs.models.encoders.llama import LlamaConfig
from fastvideo.v1.configs.models.encoders.t5 import T5Config
from fastvideo.v1.configs.models.encoders.bert import BertConfig
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig", "LlamaConfig",
"T5Config"
"T5Config", "BertConfig"
]
@@ -0,0 +1,48 @@
# fastvideo/v1/configs/models/encoders/bert.py
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
# ---------- Architecture-level hyper-parameters -----------------
@dataclass
class BertArchConfig(TextEncoderArchConfig):
vocab_size: int = 47020
hidden_size: int = 1024
intermediate_size: int = 4096
num_hidden_layers: int = 24
num_attention_heads: int = 16
hidden_act: str = "gelu"
max_position_embeddings: int = 512
type_vocab_size: int = 2
position_embedding_type: str = "absolute"
layer_norm_eps: float = 1e-12
initializer_range: float = 0.02
attention_probs_dropout_prob: float = 0.1
hidden_dropout_prob: float = 0.1
use_cache: bool = True # KV cache toggle
pad_token_id: int = 0
bos_token_id: int = 0
eos_token_id: int = 2
text_len: int = 77
# === Pooler / NSP head ===================================================
pooler_fc_size: int = 768
pooler_num_attention_heads: int = 12
pooler_num_fc_layers: int = 3
pooler_size_per_head: int = 128
pooler_type: str = "first_token_transform"
classifier_dropout: Optional[float] = None
directionality: str = "bidi"
output_past: bool = True
# ---------- Top-level config wrapper (YAML-friendly) ------------
@dataclass
class BertConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(default_factory=BertArchConfig)
prefix: str = "bert"
+2 -2
View File
@@ -146,7 +146,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
# Apply normalization
normalized = self.norm(residual_output)
# Apply scale and shift
modulated = normalized * (1.0 + scale) + shift
modulated = normalized * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
return modulated, residual_output
@@ -182,4 +182,4 @@ class LayerNormScaleShift(nn.Module):
scale: torch.Tensor) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
normalized = self.norm(x)
return normalized * (1.0 + scale) + shift
return normalized * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
+403 -32
View File
@@ -1,40 +1,411 @@
# type: ignore
import os
# SPDX-License-Identifier: Apache-2.0
from typing import Iterable, Optional, Set, Tuple
import torch
import torch.nn as nn
from transformers import BertModel, BertTokenizer
from torch import nn
from transformers import BertConfig
# from vllm.attention import Attention, AttentionType
from fastvideo.v1.attention import LocalAttention
# from vllm.compilation.decorators import support_torch_compile
# from vllm.config import CacheConfig, PoolerConfig, VllmConfig
# CacheConfig is used for attention, could be replaced with LocalAttention
# VllmConfig is replaced with BertConfig (needs to implement)
from fastvideo.v1.configs.models.encoders.bert import BertConfig
# from vllm.distributed import get_tensor_model_parallel_world_size
from fastvideo.v1.distributed import (divide,
get_tensor_model_parallel_world_size)
# from vllm.forward_context import get_forward_context
from fastvideo.v1.forward_context import get_forward_context
# from vllm.model_executor.layers.activation import get_act_fn
from fastvideo.v1.layers.activation import get_act_fn
# from vllm.model_executor.layers.linear import (ColumnParallelLinear,
# QKVParallelLinear,
# RowParallelLinear)
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
RowParallelLinear)
# from vllm.model_executor.layers.pooler import (CrossEncodingPooler, Pooler,
# PoolingType)
# from vllm.model_executor.layers.quantization import QuantizationConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
# from vllm.model_executor.layers.vocab_parallel_embedding import (
# VocabParallelEmbedding)
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
# from vllm.model_executor.model_loader.weight_utils import default_weight_loader
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
# from vllm.model_executor.pooling_metadata import PoolingMetadata
# from vllm.sequence import PoolerOutput
# from vllm.transformers_utils.config import (
# get_cross_encoder_activation_function)
from fastvideo.v1.models.encoders.base import TextEncoder
# from .interfaces import SupportsCrossEncoding, SupportsQuant, SupportsV0Only
# from .utils import WeightsMapper, maybe_prefix
from fastvideo.v1.platforms import _Backend
class BertEmbedding(nn.Module):
def __init__(self, config: BertConfig):
super().__init__()
self.size = config.hidden_size
self.word_embeddings = VocabParallelEmbedding(config.vocab_size,
config.hidden_size)
self.position_embeddings = VocabParallelEmbedding(
config.max_position_embeddings, config.hidden_size)
self.token_type_embeddings = VocabParallelEmbedding(
config.type_vocab_size, config.hidden_size)
self.LayerNorm = nn.LayerNorm(config.hidden_size,
eps=config.layer_norm_eps)
# self.position_ids = nn.Parameter(
# torch.empty((1, config.max_position_embeddings)), )
self.position_embedding_type = config.position_embedding_type
if self.position_embedding_type != "absolute":
raise ValueError("Only 'absolute' position_embedding_type" +
" is supported")
def forward(
self,
input_ids: torch.Tensor,
position_ids: torch.Tensor,
token_type_ids: Optional[torch.Tensor] = None,
) -> torch.Tensor:
input_shape = input_ids.size()
# Input embeddings.
inputs_embeds = self.word_embeddings(input_ids)
# Position embeddings.
position_embeddings = self.position_embeddings(position_ids)
if token_type_ids is None:
token_type_ids = torch.zeros(input_shape,
dtype=torch.long,
device=inputs_embeds.device)
token_type_embeddings = self.token_type_embeddings(token_type_ids)
embeddings = inputs_embeds + token_type_embeddings + position_embeddings
embeddings = self.LayerNorm(embeddings)
return embeddings
class HunyuanClip(nn.Module):
"""
Hunyuan clip code copied from https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/hunyuandit/pipeline_hunyuandit.py
hunyuan's clip used BertModel and BertTokenizer, so we copy it.
"""
class BertPooler(nn.Module):
def __init__(self, model_dir, max_length=77):
def __init__(self, config: BertConfig):
super().__init__()
self.dense = nn.Linear(config.hidden_size, config.hidden_size)
self.activation = nn.Tanh()
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
# We "pool" the model by simply taking the hidden state corresponding
# to the first token.
first_token_tensor = hidden_states[0, :]
pooled_output = self.dense(first_token_tensor)
pooled_output = self.activation(pooled_output)
return pooled_output
# @support_torch_compile
class BertEncoder(nn.Module):
def __init__(self, config: BertConfig, quant_config: Optional[QuantizationConfig] = None, prefix: str = ""):
super().__init__()
# config = vllm_config.model_config.hf_config
# cache_config = vllm_config.cache_config
# quant_config = vllm_config.quant_config
self.layer = nn.ModuleList([
BertLayer(config=config,
# cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.layer.{layer_idx}")
for layer_idx in range(config.num_hidden_layers)
])
def forward(
self,
hidden_states: torch.Tensor,
) -> torch.Tensor:
for layer in self.layer:
hidden_states = layer(hidden_states)
print("hidden_states shape", hidden_states.shape)
return hidden_states
class BertLayer(nn.Module):
def __init__(self,
config: BertConfig,
# cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
super().__init__()
self.max_length = max_length
self.tokenizer = BertTokenizer.from_pretrained(
os.path.join(model_dir, 'tokenizer'))
self.text_encoder = BertModel.from_pretrained(
os.path.join(model_dir, 'clip_text_encoder'))
self.attention = BertAttention(
hidden_size=config.hidden_size,
num_attention_heads=config.num_attention_heads,
layer_norm_eps=config.layer_norm_eps,
# cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.attention")
@torch.no_grad
def forward(self, prompts, with_mask=True):
self.device = next(self.text_encoder.parameters()).device
text_inputs = self.tokenizer(
prompts,
padding="max_length",
max_length=self.max_length,
truncation=True,
return_attention_mask=True,
return_tensors="pt",
)
prompt_embeds = self.text_encoder(
text_inputs.input_ids.to(self.device),
attention_mask=text_inputs.attention_mask.to(self.device)
if with_mask else None,
)
return prompt_embeds.last_hidden_state, prompt_embeds.pooler_output
self.intermediate = BertIntermediate(
hidden_size=config.hidden_size,
intermediate_size=config.intermediate_size,
hidden_act=config.hidden_act,
quant_config=quant_config,
prefix=f"{prefix}.intermediate")
self.output = BertOutput(hidden_size=config.hidden_size,
intermediate_size=config.intermediate_size,
layer_norm_eps=config.layer_norm_eps,
quant_config=quant_config,
prefix=f"{prefix}.output")
def forward(self, hidden_states: torch.Tensor):
attn_output = self.attention(hidden_states)
intermediate_output = self.intermediate(attn_output)
output = self.output(intermediate_output, attn_output)
return output
class BertAttention(nn.Module):
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
layer_norm_eps: float,
# cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
self.self = BertSelfAttention(hidden_size=hidden_size,
num_attention_heads=num_attention_heads,
# cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.output")
self.output = BertSelfOutput(hidden_size=hidden_size,
layer_norm_eps=layer_norm_eps,
quant_config=quant_config,
prefix=f"{prefix}.output")
def forward(
self,
hidden_states: torch.Tensor,
) -> torch.Tensor:
self_output = self.self(hidden_states)
return self.output(self_output, hidden_states)
class BertSelfAttention(nn.Module):
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
# cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size()
self.total_num_heads = num_attention_heads
assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size
self.total_num_kv_heads = self.total_num_heads
self.head_dim = self.hidden_size // self.total_num_heads
assert self.head_dim * self.total_num_heads == self.hidden_size
self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)
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.qkv_proj = QKVParallelLinear(
hidden_size=self.hidden_size,
head_size=self.head_dim,
total_num_heads=self.total_num_heads,
total_num_kv_heads=self.total_num_kv_heads,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.qkv_proj")
self.tp_size = get_tensor_model_parallel_world_size()
self.num_heads_per_partition = divide(self.num_heads, self.tp_size)
self.attn = LocalAttention(num_heads=self.num_heads,
head_size=self.head_dim,
num_kv_heads=self.num_kv_heads,
# softmax_scale=self.scaling,
# cache_config=cache_config,
# quant_config=quant_config,
# prefix=f"{prefix}.attn",
# attn_type=AttentionType.ENCODER_ONLY
causal=False,
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
) # River TODO, fix hardcoded backend
def forward(
self,
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)
B, L, _ = q.shape
q = q.reshape(B, L, self.num_heads, self.head_dim)
k = k.reshape(B, L, self.num_heads, self.head_dim)
v = v.reshape(B, L, self.num_heads, self.head_dim)
output = self.attn(q, k, v)
output = output.reshape(B, L, self.hidden_size)
return output
# def forward(
# self,
# 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)
# output = self.attn(q, k, v)
# return output
class BertSelfOutput(nn.Module):
def __init__(self,
hidden_size: int,
layer_norm_eps: float,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
super().__init__()
self.dense = RowParallelLinear(input_size=hidden_size,
output_size=hidden_size,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.dense")
self.LayerNorm = nn.LayerNorm(hidden_size, eps=layer_norm_eps)
def forward(self, hidden_states: torch.Tensor,
input_tensor: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.dense(hidden_states)
hidden_states = self.LayerNorm(hidden_states + input_tensor)
return hidden_states
class BertIntermediate(nn.Module):
def __init__(self,
hidden_size: int,
intermediate_size: int,
hidden_act: str,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
super().__init__()
self.dense = ColumnParallelLinear(input_size=hidden_size,
output_size=intermediate_size,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.dense")
self.intermediate_act_fn = get_act_fn(hidden_act)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.dense(hidden_states)
hidden_states = self.intermediate_act_fn(hidden_states)
return hidden_states
class BertOutput(nn.Module):
def __init__(self,
hidden_size: int,
intermediate_size: int,
layer_norm_eps: float,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
super().__init__()
self.dense = RowParallelLinear(input_size=intermediate_size,
output_size=hidden_size,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.dense")
self.LayerNorm = nn.LayerNorm(hidden_size, eps=layer_norm_eps)
def forward(self, hidden_states: torch.Tensor,
input_tensor: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.dense(hidden_states)
hidden_states = self.LayerNorm(hidden_states + input_tensor)
return hidden_states
class BertModel(TextEncoder):
packed_modules_mapping = {"qkv_proj": ["query", "key", "value"]}
def __init__(
self,
config: BertConfig,
) -> None:
super().__init__(config)
self.embeddings = BertEmbedding(config=config)
print("prefix here ",config.prefix)
self.encoder = BertEncoder(config=config,
prefix=f"{config.prefix}.encoder")
def forward(
self,
input_ids: torch.Tensor,
position_ids: torch.Tensor,
inputs_embeds: Optional[torch.Tensor] = None,
token_type_ids: Optional[torch.Tensor] = None,
**kwargs,
) -> torch.Tensor:
if inputs_embeds is not None:
hidden_states = inputs_embeds
else:
hidden_states = self.embeddings(
input_ids=input_ids,
position_ids=position_ids,
token_type_ids=token_type_ids)
return self.encoder(hidden_states)
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
("qkv_proj", "query", "q"),
("qkv_proj", "key", "k"),
("qkv_proj", "value", "v"),
]
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
for name, loaded_weight in weights:
name = name[len("bert."):] if name.startswith("bert.") else name
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
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 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
+2 -2
View File
@@ -36,7 +36,7 @@ _TEXT_ENCODER_MODELS = {
"LlamaModel": ("encoders", "llama", "LlamaModel"),
"UMT5EncoderModel": ("encoders", "t5", "UMT5EncoderModel"),
"STEP1TextEncoder": ("encoders", "stepllm", "STEP1TextEncoder"),
"BertModel": ("encoders", "clip", "CLIPTextModel"),
"BertModel": ("encoders", "bert", "BertModel"),
}
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
@@ -295,7 +295,7 @@ class _ModelRegistry:
model_cls = self._try_load_model_cls(arch)
if model_cls is not None:
return (model_cls, arch)
print("unsupported ",architectures)
return self._raise_for_unsupported(architectures)
-3
View File
@@ -39,9 +39,6 @@ class ParallelTiledVAE(ABC):
self.use_temporal_tiling = config.use_temporal_tiling
self.use_parallel_tiling = config.use_parallel_tiling
def to(self, device) -> 'ParallelTiledVAE':
return self
@property
def temporal_compression_ratio(self) -> int:
return cast(int, self.config.temporal_compression_ratio)
+1 -2
View File
@@ -7,7 +7,6 @@ import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.utils import PRECISION_TO_TYPE
@@ -24,7 +23,7 @@ class DecodingStage(PipelineStage):
"""
def __init__(self, vae) -> None:
self.vae: ParallelTiledVAE = vae
self.vae = vae
def forward(
self,
@@ -46,8 +46,6 @@ class EncodingStage(PipelineStage):
Returns:
The batch with encoded outputs.
"""
self.vae = self.vae.to(fastvideo_args.device)
image_path = batch.image_path
# TODO(will): remove this once we add input/output validation for stages
if image_path is None:
@@ -16,7 +16,7 @@ from huggingface_hub import hf_hub_download
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.encoders.bert import HunyuanClip # type: ignore
# from fastvideo.v1.models.encoders.bert_old import HunyuanClip # type: ignore
from fastvideo.v1.models.encoders.stepllm import STEP1TextEncoder
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
@@ -31,7 +31,7 @@ logger = init_logger(__name__)
class StepVideoPipeline(ComposedPipelineBase):
_required_config_modules = ["transformer", "scheduler", "vae"]
_required_config_modules = ["transformer", "scheduler", "vae", "tokenizer", "text_encoder"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
@@ -41,8 +41,8 @@ class StepVideoPipeline(ComposedPipelineBase):
self.add_stage(stage_name="prompt_encoding_stage",
stage=StepvideoPromptEncodingStage(
stepllm=self.get_module("text_encoder"),
clip=self.get_module("text_encoder_2"),
stepllm=self.get_module("text_encoder_2"),
clip=self.get_module("text_encoder"),
))
self.add_stage(stage_name="timestep_preparation_stage",
@@ -80,9 +80,9 @@ class StepVideoPipeline(ComposedPipelineBase):
llm_dir = os.path.join(self.model_path, "step_llm")
clip_dir = os.path.join(self.model_path, "hunyuan_clip")
text_enc = self.build_llm(llm_dir, target_device)
clip_enc = self.build_clip(clip_dir, target_device)
self.add_module("text_encoder", text_enc)
self.add_module("text_encoder_2", clip_enc)
# clip_enc = self.build_clip(clip_dir, target_device)
self.add_module("text_encoder_2", text_enc)
# self.add_module("text_encoder", clip_enc)
lib_path = (
os.path.join(
fastvideo_args.model_path,
@@ -111,7 +111,7 @@ class StepVideoPipeline(ComposedPipelineBase):
modules_config
) > 1, "model_index.json must contain at least one pipeline module"
required_modules = ["transformer", "scheduler", "vae"]
required_modules = ["transformer", "scheduler", "tokenizer", "text_encoder", "vae"]
for module_name in required_modules:
if module_name not in modules_config:
raise ValueError(
@@ -0,0 +1,158 @@
import sys
sys.argv = [sys.argv[0]]
import os
from fastvideo.v1.models.loader.component_loader import VAELoader, TokenizerLoader, TextEncoderLoader
from fastvideo.v1.fastvideo_args import FastVideoArgs
import torch
from fastvideo.v1.configs.models.vaes import StepVideoVAEConfig
from fastvideo.models.stepvideo.text_encoder.clip import HunyuanClip
from fastvideo.v1.configs.models.encoders.bert import BertConfig
import torch.distributed as dist
from fastvideo.v1.distributed.parallel_state import initialize_model_parallel
import pytest
from fastvideo.v1.forward_context import set_forward_context
from types import SimpleNamespace
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29500")
# if not dist.is_initialized():
# dist.init_process_group("gloo", rank=0, world_size=1)
# initialize_model_parallel(tensor_model_parallel_size=1)
model_path = 'cache/hub/models--FastVideo--stepvideo-t2v-diffusers/snapshots/572e0ce299de9fe2f8b843afe5afce5facb23c13'
VAE_PATH = model_path + '/vae'
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=model_path,
text_encoder_precisions=("fp16",),
text_encoder_configs=(BertConfig(),))
args.device_str = "cuda:0"
# args = FastVideoArgs(model_path=model_path, vae_precision=precision_str)
args.vae_config = StepVideoVAEConfig()
PROMPTS = [
"A beautiful sunset over the mountains",
"A futuristic city skyline at night",
"A serene beach with palm trees and clear water",
"A bustling market street in a vibrant city",
"A cozy cabin in the woods during winter"
]
@pytest.mark.usefixtures("distributed_setup")
def build_fastvideo_encoder(old_model_path):
"""Instantiate the FV tokenizer + text encoder the same way your loader will
inside the pipeline, but *outside* the full pipeline for easier unit test."""
tokenizer = TokenizerLoader().load(
model_path=os.path.join(old_model_path, "tokenizer"),
architecture="bert",
fastvideo_args=args,
)
encoder = TextEncoderLoader().load(
model_path=os.path.join(old_model_path, "text_encoder"),
architecture="bert",
fastvideo_args=args,
).to(dtype=precision, device=device)
return tokenizer, encoder
@pytest.mark.usefixtures("distributed_setup")
def test_hidden_and_pooled_equivalence():
# ----- original HF path
old_model_path = os.path.join(model_path, "hunyuan_clip")
legacy = HunyuanClip(old_model_path).to(dtype=precision, device=device).eval()
with torch.no_grad():
h_legacy, text_inputs = legacy(PROMPTS)
# ----- fastvideo path
tok, fv_enc = build_fastvideo_encoder(old_model_path)
tok_out = tok(
PROMPTS,
padding="max_length",
max_length=legacy.max_length,
truncation=True,
return_attention_mask=True,
return_tensors="pt",
)
# 2) Get the “fastvideo” tokenizer outputs
fv_tok = tok_out # already a BatchEncoding from your TokenizerLoader
# 3) Make sure they have the same keys
assert set(text_inputs.keys()) == set(fv_tok.keys()), \
f"Mismatch in keys: {text_inputs.keys()} vs {fv_tok.keys()}"
# 4) For each tensor, compare shape and contents
for key in text_inputs.keys():
a = text_inputs[key]
b = fv_tok[key]
# Move both to CPU so the comparison is unambiguous
a_cpu = a.cpu()
b_cpu = b.cpu()
# Check shapes
assert a_cpu.shape == b_cpu.shape, f"Shape mismatch for `{key}`: {a_cpu.shape} vs {b_cpu.shape}"
# Check values exactly
if a_cpu.dtype.is_floating_point:
# floats: allow exact equality since token IDs/masks are integers
assert torch.equal(a_cpu, b_cpu), f"Floating values differ for `{key}`"
print(key, torch.equal(a_cpu, b_cpu))
else:
# ints (input_ids, attention_mask): must be identical
assert torch.equal(a_cpu, b_cpu), f"Integer tensors differ for `{key}`"
print(key, torch.equal(a_cpu, b_cpu))
seq_len = tok_out.input_ids.size(1) # 77
position_ids = torch.arange(seq_len, device=device).unsqueeze(0)
position_ids = position_ids.expand(tok_out.input_ids.size(0), -1)
with set_forward_context(current_timestep=0, attn_metadata=None):
fv_out = fv_enc(
input_ids=tok_out.input_ids.to(device),
position_ids=position_ids.to(device),
attention_mask=tok_out.attention_mask.to(device),
)
h_fv= fv_out
# assume h_legacy, h_fv are Float/BFloat tensors on the same device
diff = (h_legacy - h_fv).abs()
# overall stats
max_diff = diff.max().item()
mean_diff = diff.mean().item()
print(f"[DEBUG] overall max abs diff = {max_diff:.6f}")
print(f"[DEBUG] overall mean abs diff = {mean_diff:.6f}")
# per‐position mean diff (averaged over hidden dim)
per_token = diff.mean(dim=-1) # shape [batch, seq_len]
print("[DEBUG] per-token mean diff:", per_token[0].tolist())
# if you only want to compare *valid* tokens (mask==1)
mask = tok_out.attention_mask.to(device).unsqueeze(-1) # [B, L, 1]
valid_diff = (h_legacy - h_fv).abs()[mask.bool().expand_as(h_legacy)]
print(f"[DEBUG] valid-token max diff = {valid_diff.max().item():.6f}")
print(f"[DEBUG] valid-token mean diff= {valid_diff.mean().item():.6f}")
# after you have `legacy` and `fv_enc` on the same device
legacy_state = {n: p.detach().cpu() for n, p in legacy.text_encoder.named_parameters()}
fv_state = {n: p.detach().cpu() for n, p in fv_enc.named_parameters()}
print(len(legacy_state.keys()), len(fv_state.keys()))
print([k for k in fv_state.keys() if not k.startswith('layer.')])
for name in ["embeddings.word_embeddings.weight",
"embeddings.position_embeddings.weight"]:
a = legacy_state[name]
b = fv_state[name][:47020,:]
print(f"for {name}: {a.shape}, {b.shape}")
print(f"{name} max abs diff = {(a - b).abs().max():.6f}")
for k in legacy_state.keys():
if k not in fv_state:
print(f"missing in fv_state: {k}")
# ----- 5. assertions
assert h_legacy.shape == h_fv.shape
# assert p_legacy.shape == p_fv.shape
assert torch.allclose(h_legacy, h_fv, atol=1e-4, rtol=1e-3)
# assert torch.allclose(p_legacy, p_fv)
if __name__ == "__main__":
test_hidden_and_pooled_equivalence()
-8
View File
@@ -154,14 +154,6 @@ class Worker:
logger.error(
"Worker %d in loop received KeyboardInterrupt, aborting forward pass",
self.rank)
try:
self.pipe.send(
{"error": "Operation aborted by KeyboardInterrupt"})
logger.info("Worker %d sent error response after interrupt",
self.rank)
except Exception as e:
logger.error("Worker %d failed to send error response: %s",
self.rank, str(e))
continue
-1
View File
@@ -1 +0,0 @@
__version__ = "0.1.0"