| 26 | |
| 27 | @SUBMODULES.register_module() |
| 28 | class TextEncoder(nn.Module): |
| 29 | def __init__(self, |
| 30 | pretrained_model='clip', |
| 31 | text_latent_dim=512, |
| 32 | time_embed_dim=2048, |
| 33 | dropout=0, |
| 34 | num_text_layers=4, |
| 35 | text_num_heads=4, |
| 36 | text_ff_size=2048, |
| 37 | use_text_proj=True): |
| 38 | super().__init__() |
| 39 | activation = 'gelu' |
| 40 | self.time_embed_dim = time_embed_dim |
| 41 | self.use_text_proj = use_text_proj |
| 42 | |
| 43 | if pretrained_model == 'clip': |
| 44 | self.clip, _ = clip.load('ViT-B/32', "cpu") |
| 45 | set_requires_grad(self.clip, False) |
| 46 | if text_latent_dim != 512: |
| 47 | self.text_pre_proj = nn.Linear(512, text_latent_dim) |
| 48 | else: |
| 49 | self.text_pre_proj = nn.Identity() |
| 50 | else: |
| 51 | raise NotImplementedError() |
| 52 | |
| 53 | if num_text_layers > 0: |
| 54 | self.use_text_finetune = True |
| 55 | textTransEncoderLayer = nn.TransformerEncoderLayer( |
| 56 | d_model=text_latent_dim, |
| 57 | nhead=text_num_heads, |
| 58 | dim_feedforward=text_ff_size, |
| 59 | dropout=dropout, |
| 60 | activation=activation) |
| 61 | self.textTransEncoder = nn.TransformerEncoder( |
| 62 | textTransEncoderLayer, |
| 63 | num_layers=num_text_layers) |
| 64 | else: |
| 65 | self.use_text_finetune = False |
| 66 | self.text_ln = nn.LayerNorm(text_latent_dim) |
| 67 | if self.use_text_proj: |
| 68 | self.text_proj = nn.Sequential( |
| 69 | nn.Linear(text_latent_dim, self.time_embed_dim) |
| 70 | ) |
| 71 | |
| 72 | def forward(self, text, token=None, device=None): |
| 73 | with torch.no_grad(): |
| 74 | text = clip.tokenize(text, truncate=True).to(device) |
| 75 | x = self.clip.token_embedding(text).type(self.clip.dtype) |
| 76 | |
| 77 | x = x + self.clip.positional_embedding.type(self.clip.dtype) |
| 78 | x = x.permute(1, 0, 2) |
| 79 | x = self.clip.transformer(x) |
| 80 | x = self.clip.ln_final(x).type(self.clip.dtype) |
| 81 | |
| 82 | x = self.text_pre_proj(x) |
| 83 | xf_out = self.textTransEncoder(x) |
| 84 | xf_out = self.text_ln(xf_out) |
| 85 | if self.use_text_proj: |
nothing calls this directly
no outgoing calls
no test coverage detected