File size: 3,272 Bytes
cea9f80
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
"""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()