(self, input_ids, attention_mask, token_type_ids)
| 19 | self.pooling = pooling |
| 20 | |
| 21 | def forward(self, input_ids, attention_mask, token_type_ids): |
| 22 | out = self.model(input_ids, attention_mask=attention_mask, token_type_ids=token_type_ids) |
| 23 | |
| 24 | if self.pooling == "cls": |
| 25 | return out.last_hidden_state[:, 0] |
| 26 | if self.pooling == "pooler": |
| 27 | return out.pooler_output |
| 28 | if self.pooling == 'last-avg': |
| 29 | last = out.last_hidden_state.transpose(1, 2) |
| 30 | return torch.avg_pool1d(last, kernel_size=last.shape[-1]).squeeze(-1) |
| 31 | if self.pooling == 'first-last-avg': |
| 32 | first = out.hidden_states[1].transpose(1, 2) |
| 33 | last = out.hidden_states[-1].transpose(1, 2) |
| 34 | first_avg = torch.avg_pool1d(first, kernel_size=last.shape[-1]).squeeze(-1) |
| 35 | last_avg = torch.avg_pool1d(last, kernel_size=last.shape[-1]).squeeze(-1) |
| 36 | avg = torch.cat((first_avg.unsqueeze(1), last_avg.unsqueeze(1)), dim=1) |
| 37 | return torch.avg_pool1d(avg.transpose(1, 2), kernel_size=2).squeeze(-1) |
nothing calls this directly
no outgoing calls
no test coverage detected