Team Ai
Apppublic

xdecoder/Instruct-X-Decoder

sourceHugging Faceafl-3.0updated 3y agoView on Hugging Face
163likes
vlpencoder.py168 linesDownload Raw Back to language
1 2import torch3from torch import nn4from torch.nn import functional as F5 6from timm.models.layers import trunc_normal_7 8from .registry import register_model9from ..utils import configurable10from .LangEncoder import build_tokenizer, build_lang_encoder11from utils.misc import prompt_engineering, get_prompt_templates12 13 14class LanguageEncoder(nn.Module):15 16    @configurable17    def __init__(18        self,19        tokenizer,20        tokenizer_type,21        lang_encoder,22        lang_projection,23        max_token_num,24    ):25        super().__init__()26        self.tokenizer = tokenizer27        self.tokenizer_type = tokenizer_type28        self.lang_encoder = lang_encoder29        self.lang_proj = lang_projection30        self.max_token_num = max_token_num31        self.logit_scale = nn.Parameter(torch.ones([]))32 33    @classmethod34    def from_config(cls, cfg):35        tokenizer = build_tokenizer(cfg['MODEL']['TEXT'])36        tokenizer_type = cfg['MODEL']['TEXT']['TOKENIZER']37        lang_encoder = build_lang_encoder(cfg['MODEL']['TEXT'], tokenizer, cfg['VERBOSE'])38        max_token_num = cfg['MODEL']['TEXT']['CONTEXT_LENGTH']39        40        dim_lang = cfg['MODEL']['TEXT']['WIDTH']41        dim_projection = cfg['MODEL']['DIM_PROJ']42        lang_projection = nn.Parameter(torch.empty(dim_lang, dim_projection))43        trunc_normal_(lang_projection, std=.02)44        45        return {46            "tokenizer": tokenizer,47            "tokenizer_type": tokenizer_type,48            "lang_encoder": lang_encoder,49            "lang_projection": lang_projection,50            "max_token_num": max_token_num,51        }52 53    def get_text_embeddings(self, class_names, name='default', is_eval=False, add_bgd=False, prompt=True, norm=True):54        if not is_eval:55            if prompt:56                # randomly sample one template57                arbitary_concepts = [58                    prompt_engineering(class_names[label].replace('-other','').replace('-merged','').replace('-stuff',''), topk=10000, suffix='.') \59                    for label in range(len(class_names))60                ]61                if add_bgd:62                    arbitary_concepts.append("A background in coco.")63            else:64                arbitary_concepts = class_names65            66            input_ids = []67            attention_masks = []68            for txt in arbitary_concepts:69                tokens = self.tokenizer(70                    txt, padding='max_length', truncation=True, max_length=self.max_token_num, return_tensors='pt'71                )72                tokens['input_ids'].squeeze_()73                tokens['attention_mask'].squeeze_()74 75                input_ids.append(tokens['input_ids'])76                attention_masks.append(tokens['attention_mask'])77 78            arbitary_tokens = torch.stack(input_ids)79            arbitary_attention_masks = torch.stack(attention_masks)80 81            text_emb = self.forward_language((arbitary_tokens.cuda(), arbitary_attention_masks.cuda()), norm=norm)82            setattr(self, '{}_text_embeddings'.format(name), text_emb)83        else:84            with torch.no_grad():85                def extract_mean_emb(txts):86                    tokens = self.tokenizer(87                        txts, padding='max_length', truncation=True, max_length=self.max_token_num, return_tensors='pt'88                    )89                    clss_embedding = self.forward_language((tokens['input_ids'].cuda(), tokens['attention_mask'].cuda()), norm=norm)90                    clss_embedding = clss_embedding.mean(dim=0)91                    clss_embedding /= clss_embedding.norm()92                    return clss_embedding93 94                templates = get_prompt_templates()95                clss_embeddings = []96                if prompt:97                    for clss in class_names:98                        txts = [template.format(clss.replace('-other','').replace('-merged','').replace('-stuff','')) for template in templates]99                        clss_embeddings.append(extract_mean_emb(txts))100                else:101                    clss_embeddings.append(extract_mean_emb(class_names))102 103                if add_bgd:104                    txts = ["A background in coco."]105                    clss_embeddings.append(extract_mean_emb(txts))106 107                text_emb = torch.stack(clss_embeddings, dim=0)108                setattr(self, '{}_text_embeddings'.format(name), text_emb)109 110    def get_text_token_embeddings(self, txts, name='default', token=False, norm=False):111        if not token:112            tokens = self.tokenizer(113                txts, padding='max_length', truncation=True, max_length=self.max_token_num, return_tensors='pt'114            )115            tokens = {key: value.cuda() for key, value in tokens.items()}116        else:117            tokens = txts118        token_emb, class_emb = self.forward_language_token((tokens['input_ids'], tokens['attention_mask']), norm=norm)119        ret = {"tokens": tokens,120                "token_emb": token_emb,121                "class_emb": class_emb,}122        setattr(self, '{}_token_embeddings'.format(name), ret)123        return ret124 125    def forward_language(self, texts, norm=True):126        x = self.lang_encoder(*texts)127        x = x['last_hidden_state']128 129        if self.tokenizer_type == 'clip':130            x = x[torch.arange(x.size(0)), texts[0].argmax(dim=-1)]131        else:132            x = x[:, 0]133 134        x = x @ self.lang_proj135        if norm:136            x = x / (x.norm(dim=-1, keepdim=True) + 1e-7)137        return x138    139    def forward_language_token(self, texts, norm=False):140        x = self.lang_encoder(*texts)141        token_x = x['last_hidden_state']142 143        if self.tokenizer_type == 'clip':144            class_x = token_x[torch.arange(token_x.size(0)), texts[0].argmax(dim=-1)]145        else:146            class_x = token_x[:, 0]147 148        class_x = class_x @ self.lang_proj149        token_x = token_x @ self.lang_proj150 151        if norm:152            class_x = class_x / (class_x.norm(dim=-1, keepdim=True) + 1e-7)153            token_x = token_x / (token_x.norm(dim=-1, keepdim=True) + 1e-7)154 155        return token_x, class_x156    157    def compute_similarity(self, v_emb, name='default', fake=False):158        if fake:159            return None160        v_emb = v_emb / (v_emb.norm(dim=-1, keepdim=True) + 1e-7)161        t_emb = getattr(self, '{}_text_embeddings'.format(name))162        output = self.logit_scale.exp() * v_emb @ t_emb.unsqueeze(0).transpose(1, 2)163        return output164 165 166@register_model167def get_language_model(cfg, **kwargs):168    return LanguageEncoder(cfg)