MCPcopy Create free account
hub / github.com/ZBayes/basic_rag / SimcseModel

Class SimcseModel

src/models/vec_model/simcse_model.py:7–37  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5from transformers import BertConfig, BertModel, BertTokenizer
6
7class 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)

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected