(cls, cfg)
| 31 | |
| 32 | @classmethod |
| 33 | def from_config(cls, cfg): |
| 34 | # build up text encoder |
| 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 | # @torch.no_grad() |
| 54 | def get_text_embeddings(self, class_names, name='default', is_eval=False, add_bgd=False, prompt=True, norm=True): |
nothing calls this directly
no test coverage detected