google/tipsv1-s14
2367
1"""TIPSv2 model for HuggingFace — wraps vision and text encoders."""2 3from dataclasses import dataclass4from typing import List, Optional, Union5 6import torch7from transformers import PreTrainedModel8from transformers.utils import cached_file9 10from .configuration_tips import TIPSv2Config11from .image_encoder import vit_base, vit_giant2, vit_large, vit_small, vit_so400m12from .text_encoder import TextEncoder, Tokenizer13 14_VISION_FACTORIES = {15 "vit_small": vit_small,16 "vit_base": vit_base,17 "vit_large": vit_large,18 "vit_so400m": vit_so400m,19 "vit_giant2": vit_giant2,20}21 22 23@dataclass24class TIPSv2ImageOutput:25 """Output from the vision encoder."""26 cls_token: torch.Tensor # (B, 1, D)27 register_tokens: torch.Tensor # (B, R, D)28 patch_tokens: torch.Tensor # (B, N, D)29 30 31@dataclass32class TIPSv2Output:33 """Output from the full model."""34 image_features: Optional[TIPSv2ImageOutput] = None35 text_embeds: Optional[torch.Tensor] = None36 temperature: Optional[float] = None37 38 39class TIPSv2Model(PreTrainedModel):40 """TIPSv2 vision-language model.41 42 Usage::43 44 model = AutoModel.from_pretrained("google/tipsv2-b14", trust_remote_code=True)45 46 # Image features47 out = model.encode_image(pixel_values) # pixel_values in [0, 1]48 cls = out.cls_token # (B, 1, D)49 spatial = out.patch_tokens # (B, N, D)50 51 # Text features52 text_emb = model.encode_text(["a photo of a cat"]) # (B, D)53 """54 55 config_class = TIPSv2Config56 _no_split_modules = []57 _supports_cache_class = False58 _tied_weights_keys = []59 60 @property61 def all_tied_weights_keys(self):62 return {}63 64 def __init__(self, config: TIPSv2Config):65 super().__init__(config)66 67 self.vision_encoder = _VISION_FACTORIES[config.vision_fn](68 img_size=config.img_size,69 patch_size=config.patch_size,70 ffn_layer=config.ffn_layer,71 block_chunks=0,72 init_values=config.init_values,73 interpolate_antialias=True,74 interpolate_offset=0.0,75 )76 77 self.text_encoder = TextEncoder(78 config={79 "hidden_size": config.text_hidden_size,80 "mlp_dim": config.text_mlp_dim,81 "num_heads": config.text_num_heads,82 "num_layers": config.text_num_layers,83 },84 vocab_size=config.vocab_size,85 )86 87 self._tokenizer = None88 89 def _load_tokenizer(self):90 """Load the SentencePiece tokenizer shipped with the checkpoint."""91 return Tokenizer(cached_file(self.name_or_path, "tokenizer.model"))92 93 @torch.no_grad()94 def encode_image(self, pixel_values: torch.Tensor) -> TIPSv2ImageOutput:95 """Encode images. pixel_values: (B, 3, H, W) in [0, 1]."""96 pixel_values = pixel_values.to(self.device)97 cls_token, register_tokens, patch_tokens = self.vision_encoder(pixel_values)98 return TIPSv2ImageOutput(99 cls_token=cls_token,100 register_tokens=register_tokens,101 patch_tokens=patch_tokens,102 )103 104 @torch.no_grad()105 def encode_text(106 self,107 texts: Union[str, List[str], torch.Tensor],108 padding_mask: Optional[torch.Tensor] = None,109 ) -> torch.Tensor:110 """Encode text. Pass strings (auto-tokenized) or pre-tokenized tensors."""111 if isinstance(texts, (str, list)):112 if isinstance(texts, str):113 texts = [texts]114 if self._tokenizer is None:115 self._tokenizer = self._load_tokenizer()116 ids, paddings = self._tokenizer.tokenize(texts, max_len=self.config.max_len)117 ids = torch.from_numpy(ids).to(self.device)118 padding_mask = torch.from_numpy(paddings).to(self.device)119 else:120 ids = texts.to(self.device)121 padding_mask = padding_mask.to(self.device)122 return self.text_encoder(ids, padding_mask)123 124 def forward(125 self,126 pixel_values: Optional[torch.Tensor] = None,127 input_ids: Optional[torch.Tensor] = None,128 padding_mask: Optional[torch.Tensor] = None,129 ) -> TIPSv2Output:130 """Forward pass for both or either modality."""131 image_features = None132 text_embeds = None133 if pixel_values is not None:134 image_features = self.encode_image(pixel_values)135 if input_ids is not None:136 text_embeds = self.encode_text(input_ids, padding_mask)137 return TIPSv2Output(138 image_features=image_features,139 text_embeds=text_embeds,140 temperature=self.config.temperature,141 )142 