Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
702e77c937 | ||
|
|
2d43f78d04 | ||
|
|
e21a124f06 |
@@ -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
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user