| 17 | |
| 18 | |
| 19 | class LanguageEncoder(nn.Module): |
| 20 | |
| 21 | @configurable |
| 22 | def __init__( |
| 23 | self, |
| 24 | tokenizer, |
| 25 | tokenizer_type, |
| 26 | lang_encoder, |
| 27 | lang_projection, |
| 28 | max_token_num, |
| 29 | queue_operator, |
| 30 | ): |
| 31 | super().__init__() |
| 32 | # seg |
| 33 | self.tokenizer = tokenizer |
| 34 | self.tokenizer_type = tokenizer_type |
| 35 | self.lang_encoder = lang_encoder |
| 36 | self.lang_proj = lang_projection |
| 37 | self.max_token_num = max_token_num |
| 38 | self.logit_scale = nn.Parameter(torch.ones([])) |
| 39 | |
| 40 | # captioning & retrieval |
| 41 | for key, value in queue_operator.items(): |
| 42 | self.register_buffer(key, value) |
| 43 | |
| 44 | |
| 45 | @classmethod |
| 46 | def from_config(cls, cfg): |
| 47 | # build up text encoder for seg |
| 48 | tokenizer = build_tokenizer(cfg['MODEL']['TEXT']) |
| 49 | tokenizer_type = cfg['MODEL']['TEXT']['TOKENIZER'] |
| 50 | lang_encoder = build_lang_encoder(cfg['MODEL']['TEXT'], tokenizer, cfg['VERBOSE']) |
| 51 | max_token_num = cfg['MODEL']['TEXT']['CONTEXT_LENGTH'] |
| 52 | |
| 53 | dim_lang = cfg['MODEL']['TEXT']['WIDTH'] |
| 54 | dim_projection = cfg['MODEL']['DIM_PROJ'] |
| 55 | lang_projection = nn.Parameter(torch.empty(dim_lang, dim_projection)) |
| 56 | trunc_normal_(lang_projection, std=.02) |
| 57 | |
| 58 | # tested not working better |
| 59 | queue_operator = {} |
| 60 | |
| 61 | return { |
| 62 | "tokenizer": tokenizer, |
| 63 | "tokenizer_type": tokenizer_type, |
| 64 | "lang_encoder": lang_encoder, |
| 65 | "lang_projection": lang_projection, |
| 66 | "max_token_num": max_token_num, |
| 67 | "queue_operator": queue_operator, |
| 68 | } |
| 69 | |
| 70 | def get_text_embeddings(self, class_names, name='default', is_eval=False, add_bgd=False, prompt=True, norm=True): |
| 71 | if not is_eval: |
| 72 | if prompt: |
| 73 | # randomly sample one template |
| 74 | arbitary_concepts = [ |
| 75 | prompt_engineering(class_names[label].replace('-other','').replace('-merged','').replace('-stuff',''), topk=10000, suffix='.') \ |
| 76 | for label in range(len(class_names)) |
no outgoing calls
no test coverage detected