fuse / modeling_fuse.py
harpertoken's picture
Upload modeling_fuse.py with huggingface_hub
cea9f80 verified
Raw History Blame Contribute Delete
3.27 kB
"""Transformers-compatible wrapper for harpertoken/fuse.
The model matches images to captions with a dot product between projected
CNN and embedding representations, not a transformer. This subclasses
PreTrainedModel so the weights serialize honestly as config.json plus
model.safetensors and load through AutoModel. Importing this module registers
the architecture.
"""
import torch
import torch.nn as nn
from transformers import AutoConfig, AutoModel, PretrainedConfig, PreTrainedModel
class FuseConfig(PretrainedConfig):
model_type = "fuse_matcher"
def __init__(self, vocab_size=4433, num_labels=1, dropout=0.25, **kwargs):
super().__init__(**kwargs)
self.vocab_size = vocab_size
self.num_labels = num_labels
self.dropout = dropout
class FuseMatcher(PreTrainedModel):
config_class = FuseConfig
base_model_prefix = "fuse"
# Transformers 5 reads this mapping while finalising a load. Nothing here
# is tied, but the attribute has to exist or from_pretrained raises
# AttributeError before the weights are placed.
all_tied_weights_keys = {}
def __init__(self, config):
super().__init__(config)
self.cnn = nn.ModuleDict(
{
"conv1": nn.Conv2d(3, 32, 3, padding=1),
"conv2": nn.Conv2d(32, 64, 3, padding=1),
"fc1": nn.Linear(64 * 8 * 8, 128),
}
)
self.pool = nn.MaxPool2d(2, 2)
self.adapt = nn.AdaptiveAvgPool2d((8, 8))
self.emb = nn.Embedding(config.vocab_size, 64, padding_idx=0)
self.tfc = nn.Linear(64, 128)
self.iproj = nn.Linear(128, 128)
self.tproj = nn.Linear(128, 128)
self.logit_scale = nn.Parameter(torch.tensor(4.0))
self.drop = nn.Dropout(config.dropout)
def forward(self, image, tokens):
import torch.nn.functional as F
x = self.pool(torch.relu(self.cnn["conv1"](image)))
x = self.pool(torch.relu(self.cnn["conv2"](x)))
x = self.adapt(x)
xi = self.iproj(self.drop(torch.relu(self.cnn["fc1"](x.view(x.size(0), -1)))))
mask = (tokens != 0).float().unsqueeze(-1)
te = self.emb(tokens) * mask
xt = self.tproj(
self.drop(torch.relu(self.tfc(te.sum(1) / mask.sum(1).clamp(min=1))))
)
xi = F.normalize(xi, dim=-1)
xt = F.normalize(xt, dim=-1)
return self.logit_scale.exp().clamp(max=100.0) * (xi * xt).sum(-1)
def match(self, image, caption, vocab):
"""image: HxWx3 float array in [0, 1]. caption: string.
Returns (match_probability, predicted_match)."""
import re
self.eval()
img = torch.as_tensor(image, dtype=torch.float32).permute(2, 0, 1)
ids = [vocab.get(w, 1) for w in re.findall(r"[a-z']+", caption.lower())]
ids = ids[:32] + [0] * (32 - len(ids))
with torch.no_grad():
logit = self.forward(
img.unsqueeze(0),
torch.tensor([ids]),
)[0]
prob = float(torch.sigmoid(logit))
return prob, prob > 0.5
def register():
AutoConfig.register(FuseConfig.model_type, FuseConfig, exist_ok=True)
AutoModel.register(FuseConfig, FuseMatcher, exist_ok=True)
register()