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
|