| 7 | """Text Encoder""" |
| 8 | |
| 9 | class CLIPTextEncoder(nn.Module): |
| 10 | def __init__(self, context_length=77, |
| 11 | vocab_size=49408, |
| 12 | # vocab_size=49408+1, |
| 13 | transformer_width=512, |
| 14 | transformer_heads=8, |
| 15 | transformer_layers=12, |
| 16 | embed_dim=512, |
| 17 | out_dim=256, |
| 18 | pretrained=None, **kwargs): |
| 19 | super().__init__() |
| 20 | |
| 21 | self.pretrained = pretrained |
| 22 | |
| 23 | self.context_length = context_length |
| 24 | |
| 25 | self.transformer = Transformer( |
| 26 | width=transformer_width, |
| 27 | layers=transformer_layers, |
| 28 | heads=transformer_heads, |
| 29 | attn_mask=self.build_attention_mask() |
| 30 | ) |
| 31 | |
| 32 | self.vocab_size = vocab_size |
| 33 | self.token_embedding = nn.Embedding(vocab_size, transformer_width) |
| 34 | self.positional_embedding = nn.Parameter(torch.empty(self.context_length, transformer_width)) |
| 35 | self.ln_final = LayerNorm(transformer_width) |
| 36 | self.text_projection = nn.Parameter(torch.empty(transformer_width, embed_dim)) |
| 37 | # self.text_projection = nn.Linear(transformer_width, embed_dim) |
| 38 | |
| 39 | def init_weights(self, pretrained=None): |
| 40 | pretrained = pretrained or self.pretrained |
| 41 | if isinstance(pretrained, str): |
| 42 | checkpoint = torch.jit.load(pretrained, map_location='cpu').float().state_dict() |
| 43 | # checkpoint = torch.load(pretrained)['model'] |
| 44 | |
| 45 | state_dict = {} |
| 46 | |
| 47 | for k in checkpoint.keys(): |
| 48 | if k.startswith('transformer.'): |
| 49 | # if k.startswith('module.encode_text.transformer.'): |
| 50 | # new_k = k.replace('module.encode_text.', '') |
| 51 | # state_dict[new_k] = checkpoint[k].float() |
| 52 | state_dict[k] = checkpoint[k].float() |
| 53 | |
| 54 | if k == 'positional_embedding' or k == 'text_projection' or k.startswith('token_embedding') or k.startswith('ln_final'): |
| 55 | # if k == 'module.encode_text.positional_embedding' or k.startswith('module.encode_text.text_projection') or k.startswith('module.encode_text.token_embedding') or k.startswith('module.encode_text.ln_final'): |
| 56 | # new_k = k.replace('module.encode_text.', '') |
| 57 | # if new_k == 'positional_embedding' and checkpoint[k].size(0) > self.context_length: |
| 58 | if k == 'positional_embedding' and checkpoint[k].size(0) > self.context_length: |
| 59 | checkpoint[k] = checkpoint[k][:self.context_length] |
| 60 | print('positional_embedding is tuncated from 77 to', self.context_length) |
| 61 | # state_dict[new_k] = checkpoint[k] |
| 62 | state_dict[k] = checkpoint[k] |
| 63 | |
| 64 | u, w = self.load_state_dict(state_dict, False) |
| 65 | if u != [] or w != [] : |
| 66 | print(u, w, 'are misaligned params in text encoder') |