| 8 | |
| 9 | |
| 10 | class ConditionEncoder(nn.Module): |
| 11 | |
| 12 | def train(self, mode: bool = True): |
| 13 | super().train(mode) |
| 14 | self.multi_encoder.eval() |
| 15 | for param in self.multi_encoder.parameters(): |
| 16 | param.requires_grad = False |
| 17 | return self |
| 18 | |
| 19 | def __init__(self, args, hidden_size): |
| 20 | super().__init__() |
| 21 | self.n_learnable_queries = 32 |
| 22 | config = AutoConfig.from_pretrained(args.text_condition) |
| 23 | self.multi_encoder = AutoModel.from_config(config) |
| 24 | qformer_config = Blip2QFormerConfig( |
| 25 | num_hidden_layers=12, |
| 26 | encoder_hidden_size=self.multi_encoder.config.text_config.hidden_size |
| 27 | ) |
| 28 | self.qformer = Blip2QFormerModel(qformer_config) |
| 29 | self.query_embeds = nn.Embedding(self.n_learnable_queries, qformer_config.hidden_size) |
| 30 | self.out_project = nn.Linear(qformer_config.hidden_size, hidden_size) |
| 31 | |
| 32 | @torch.no_grad() |
| 33 | def encode_text(self, input_ids, attention_mask): |
| 34 | text_encoder_output = self.multi_encoder.text_model( |
| 35 | input_ids=input_ids, |
| 36 | attention_mask=attention_mask |
| 37 | ) |
| 38 | text_embeds = text_encoder_output.last_hidden_state |
| 39 | return text_embeds # bs x ntoken x ch |
| 40 | |
| 41 | def forward(self, input_ids, attention_mask): |
| 42 | net_device = next(self.parameters()).device |
| 43 | text_embeds = self.encode_text(input_ids=input_ids, attention_mask=attention_mask) |
| 44 | query_embeds = self.query_embeds( |
| 45 | repeat( |
| 46 | torch.arange(0, self.n_learnable_queries, dtype=torch.int64).to(net_device), |
| 47 | 'src -> bs src', |
| 48 | bs = text_embeds.shape[0] |
| 49 | ) |
| 50 | ) |
| 51 | query_outputs = self.qformer( |
| 52 | query_embeds=query_embeds, |
| 53 | encoder_hidden_states=text_embeds, |
| 54 | encoder_attention_mask=attention_mask |
| 55 | ) |
| 56 | query_outputs = query_outputs[0][:, : self.n_learnable_queries, :] |
| 57 | return self.out_project(query_outputs) |
| 58 | |
| 59 | |
| 60 | |