Compare commits

...
3 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
7 changed files with 624 additions and 46 deletions
@@ -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"
+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)
@@ -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()