MCPcopy Create free account
hub / github.com/Seunggu0305/VLCounter / CLIPTextEncoder

Class CLIPTextEncoder

tools/models/Text_Encoder.py:9–86  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

7"""Text Encoder"""
8
9class 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')

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected