File size: 4,166 Bytes
4e2a1b3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
import torch
from PIL import Image
from torchvision import transforms
from transformers import CLIPModel as HFCLIPModel, CLIPConfig


class BioCLIPv2Model(HFCLIPModel):
    def __init__(self):
        super().__init__(self._build_config())

    @staticmethod
    def _build_config():
        return CLIPConfig(
            projection_dim=768,
            logit_scale_init_value=2.6592,
            text_config={
                "hidden_size": 768,
                "intermediate_size": 3072,
                "num_attention_heads": 12,
                "num_hidden_layers": 12,
                "max_position_embeddings": 77,
                "vocab_size": 49408,
                "hidden_act": "gelu",
                "layer_norm_eps": 1e-5,
                "projection_dim": 768,
                "bos_token_id": 0,
                "eos_token_id": 2,
                "pad_token_id": 1,
            },
            vision_config={
                "hidden_size": 1024,
                "intermediate_size": 4096,
                "num_attention_heads": 16,
                "num_hidden_layers": 24,
                "image_size": 224,
                "patch_size": 14,
                "hidden_act": "gelu",
                "layer_norm_eps": 1e-5,
                "projection_dim": 768,
            },
        )


class BioCLIPv2Compute(torch.nn.Module):
    MEAN = (0.48145466, 0.4578275, 0.40821073)
    STD = (0.26862954, 0.26130258, 0.27577711)

    def __init__(self, model: BioCLIPv2Model, tokenizer, max_length: int = 77):
        super().__init__()
        self.model = model
        self.tokenizer = tokenizer
        self.max_length = max_length
        self.image_transform = transforms.Compose([
            transforms.Resize(224, interpolation=transforms.InterpolationMode.BICUBIC),
            transforms.CenterCrop(224),
            transforms.ToTensor(),
            transforms.Normalize(self.MEAN, self.STD),
        ])

    @property
    def device(self):
        return next(self.model.parameters()).device

    @property
    def dtype(self):
        return next(self.model.parameters()).dtype

    def _preprocess_images(self, images):
        if isinstance(images, Image.Image):
            images = [images]
        images = [img.convert("RGB") for img in images]
        pixel_values = torch.stack([self.image_transform(img) for img in images])
        return pixel_values.to(device=self.device, dtype=self.dtype)

    def _tokenize(self, text):
        if isinstance(text, str):
            text = [text]
        tokens = self.tokenizer(
            text, padding=True, truncation=True,
            max_length=self.max_length, return_tensors="pt",
        )
        return {k: v.to(self.device) for k, v in tokens.items()}

    @torch.no_grad()
    def get_image_features(self, images):
        pixel_values = self._preprocess_images(images)
        features = self.model.get_image_features(pixel_values=pixel_values)
        return torch.nn.functional.normalize(features, dim=-1)

    @torch.no_grad()
    def get_text_features(self, text):
        tokens = self._tokenize(text)
        features = self.model.get_text_features(**tokens)
        return torch.nn.functional.normalize(features, dim=-1)

    @torch.no_grad()
    def forward(self, text: str | list[str], images):
        if isinstance(text, str):
            text = [text]
        if isinstance(images, Image.Image):
            images = [images]
        if len(text) == 1 and len(images) > 1:
            text = text * len(images)
        if len(images) == 1 and len(text) > 1:
            images = images * len(text)
        image_features = self.get_image_features(images)
        text_features = self.get_text_features(text)
        scores = (text_features * image_features).sum(dim=-1)
        scores = self.model.logit_scale.exp() * scores
        return scores

    @torch.no_grad()
    def similarity_matrix(self, text: str | list[str], images):
        image_features = self.get_image_features(images)
        text_features = self.get_text_features(text)
        scores = text_features @ image_features.T
        scores = self.model.logit_scale.exp() * scores
        return scores