ouzhang57's picture
Upload folder using huggingface_hub (part 10)
4e2a1b3 verified
Raw
History Blame Contribute Delete
4.17 kB
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