| 5 | from transformers import BertConfig, BertModel, BertTokenizer |
| 6 | |
| 7 | class SimcseModel(nn.Module): |
| 8 | # https://blog.csdn.net/qq_44193969/article/details/126981581 |
| 9 | def __init__(self, pretrained_bert_path, pooling="cls") -> None: |
| 10 | super(SimcseModel, self).__init__() |
| 11 | |
| 12 | self.pretrained_bert_path = pretrained_bert_path |
| 13 | self.config = BertConfig.from_pretrained(self.pretrained_bert_path) |
| 14 | |
| 15 | self.model = BertModel.from_pretrained(self.pretrained_bert_path, config=self.config) |
| 16 | self.model.eval() |
| 17 | |
| 18 | # self.model = None |
| 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) |