"""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()