xdecoder/Instruct-X-Decoder
163
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)