shadowlilac's picture
Add new SentenceTransformer model
9122db9 verified
Raw
History Blame Contribute Delete
5.17 kB
import os
import torch
from torch import nn
try:
from sentence_transformers.sentence_transformer.modules import Module
except ImportError: # older sentence-transformers layouts
try:
from sentence_transformers.base.modules import Module
except ImportError:
from sentence_transformers.models.Module import Module
class SiglipStyleMLP(nn.Module):
"""Mirrors Siglip2MLP: fc1 -> gelu_pytorch_tanh -> fc2."""
def __init__(self, hidden_size: int, intermediate_size: int) -> None:
super().__init__()
self.fc1 = nn.Linear(hidden_size, intermediate_size)
self.activation_fn = nn.GELU(approximate="tanh")
self.fc2 = nn.Linear(intermediate_size, hidden_size)
def forward(self, hidden_state: torch.Tensor) -> torch.Tensor:
return self.fc2(self.activation_fn(self.fc1(hidden_state)))
class MultiheadAttentionPooling(Module):
"""Multihead Attention Pooling, replicating Siglip2MultiheadAttentionPoolingHead.
A learned probe token attends over the token embeddings via nn.MultiheadAttention,
followed by LayerNorm and a residual MLP. The final sentence embedding is the
(single) probe position of the output: hidden_state[:, 0].
"""
config_keys: list = ["hidden_size", "num_attention_heads", "intermediate_size", "layer_norm_eps"]
def __init__(
self,
hidden_size: int,
num_attention_heads: int = 8,
intermediate_size: int | None = None,
layer_norm_eps: float = 1e-6,
**kwargs,
) -> None:
super().__init__()
if intermediate_size is None:
intermediate_size = 4 * hidden_size
assert hidden_size % num_attention_heads == 0, "hidden_size must be divisible by num_attention_heads"
self.hidden_size = hidden_size
self.num_attention_heads = num_attention_heads
self.intermediate_size = intermediate_size
self.layer_norm_eps = layer_norm_eps
self.probe = nn.Parameter(torch.randn(1, 1, hidden_size))
self.attention = torch.nn.MultiheadAttention(hidden_size, num_attention_heads, batch_first=True)
self.layernorm = nn.LayerNorm(hidden_size, eps=layer_norm_eps)
self.mlp = SiglipStyleMLP(hidden_size, intermediate_size)
self.num_heads = num_attention_heads
def forward(self, features: dict, **kwargs) -> dict:
hidden_state = features["token_embeddings"]
attention_mask = features.get("attention_mask", None)
batch_size = hidden_state.shape[0]
probe = self.probe.to(hidden_state.dtype).repeat(batch_size, 1, 1)
attn_mask = None
if attention_mask is not None:
target_len, source_len = probe.shape[1], hidden_state.shape[1]
# Equivalent of create_bidirectional_mask for this cross attention:
# expand [batch, source_len] -> [batch, 1, target_len, source_len], True = attend.
mask = attention_mask.to(torch.bool)[:, None, None, :].expand(batch_size, 1, target_len, source_len)
# Exactly as in Siglip2MultiheadAttentionPoolingHead:
mask = mask.repeat(1, self.num_heads, 1, 1)
mask = mask.reshape(-1, target_len, source_len)
# nn.MultiheadAttention cannot handle boolean masks (which SDPA can)
attn_mask = torch.where(
mask,
torch.full((), 0.0, device=mask.device, dtype=probe.dtype),
torch.finfo(probe.dtype).min,
)
hidden_state = self.attention(probe, hidden_state, hidden_state, attn_mask=attn_mask)[0]
residual = hidden_state
hidden_state = self.layernorm(hidden_state)
hidden_state = residual + self.mlp(hidden_state)
features["sentence_embedding"] = hidden_state[:, 0]
return features
def get_embedding_dimension(self) -> int:
return self.hidden_size
def save(self, output_path: str, *args, safe_serialization: bool = True, **kwargs) -> None:
self.save_config(output_path)
if safe_serialization:
from safetensors.torch import save_model
save_model(self, os.path.join(output_path, "model.safetensors"))
else:
torch.save(self.state_dict(), os.path.join(output_path, "pytorch_model.bin"))
@classmethod
def load(cls, model_name_or_path: str, subfolder: str = "", **kwargs):
hub_kwargs = {
k: kwargs[k]
for k in ("token", "cache_folder", "revision", "local_files_only")
if k in kwargs
}
config = cls.load_config(model_name_or_path=model_name_or_path, subfolder=subfolder, **hub_kwargs)
module = cls(**config)
try:
weights_path = cls.load_file_path(
model_name_or_path, filename="model.safetensors", subfolder=subfolder, **hub_kwargs
)
if weights_path:
from safetensors.torch import load_file
module.load_state_dict(load_file(weights_path))
except Exception as exc:
print(f"[MultiheadAttentionPooling] no saved weights loaded ({exc}), using fresh initialization")
return module