CoolFace
Modelpublic

google/tipsv1-s14

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
2likes367downloads
modeling_tips.py142 linesDownload Raw Back to root
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