Views
No views yet
pip install -U sentence-transformers1from sentence_transformers import SentenceTransformer
2sentences = ["This is an example sentence", "Each sentence is converted"]
3
4model = SentenceTransformer('sentence-transformers/all-mpnet-base-v2')
5embeddings = model.encode(sentences)
6print(embeddings)1from transformers import AutoTokenizer, AutoModel
2import torch
3import torch.nn.functional as F
4
5#Mean Pooling - Take attention mask into account for correct averaging
6def mean_pooling(model_output, attention_mask):
7 token_embeddings = model_output[0] #First element of model_output contains all token embeddings
8 input_mask_expanded = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
9 return torch.sum(token_embeddings * input_mask_expanded, 1) / torch.clamp(input_mask_expanded.sum(1), min=1e-9)
10
11
12# Sentences we want sentence embeddings for
13sentences = ['This is an example sentence', 'Each sentence is converted']
14
15# Load model from HuggingFace Hub
16tokenizer = AutoTokenizer.from_pretrained('sentence-transformers/all-mpnet-base-v2')
17model = AutoModel.from_pretrained('sentence-transformers/all-mpnet-base-v2')
18
19# Tokenize sentences
20encoded_input = tokenizer(sentences, padding=True, truncation=True, return_tensors='pt')
21
22# Compute token embeddings
23with torch.no_grad():
24 model_output = model(**encoded_input)
25
26# Perform pooling
27sentence_embeddings = mean_pooling(model_output, encoded_input['attention_mask'])
28
29# Normalize embeddings
30sentence_embeddings = F.normalize(sentence_embeddings, p=2, dim=1)
31
32print("Sentence embeddings:")
33print(sentence_embeddings)microsoft/mpnet-base model and fine-tuned in on a
1B sentence pairs dataset. We use a contrastive learning objective: given a sentence from the pair, the model should predict which out of a set of randomly sampled other sentences, was actually paired with it in our dataset.microsoft/mpnet-base model. Please refer to the model card for more detailed information about the pre-training procedure.train_script.py.data_config.json file.| Dataset | Paper | Number of training tuples |
|---|---|---|
| Reddit comments (2015-2018) | paper | 726,484,430 |
| S2ORC Citation pairs (Abstracts) | paper | 116,288,806 |
| WikiAnswers Duplicate question pairs | paper | 77,427,422 |
| PAQ (Question, Answer) pairs | paper | 64,371,441 |
| S2ORC Citation pairs (Titles) | paper | 52,603,982 |
| S2ORC (Title, Abstract) | paper | 41,769,185 |
| Stack Exchange (Title, Body) pairs | - | 25,316,456 |
| Stack Exchange (Title+Body, Answer) pairs | - | 21,396,559 |
| Stack Exchange (Title, Answer) pairs | - | 21,396,559 |
| MS MARCO triplets | paper | 9,144,553 |
| GOOAQ: Open Question Answering with Diverse Answer Types | paper | 3,012,496 |
| Yahoo Answers (Title, Answer) | paper | 1,198,260 |
| Code Search | - | 1,151,414 |
| COCO Image captions | paper | 828,395 |
| SPECTER citation triplets | paper | 684,100 |
| Yahoo Answers (Question, Answer) | paper | 681,164 |
| Yahoo Answers (Title, Question) | paper | 659,896 |
| SearchQA | paper | 582,261 |
| Eli5 | paper | 325,475 |
| Flickr 30k | paper | 317,695 |
| Stack Exchange Duplicate questions (titles) | 304,525 | |
| AllNLI (SNLI and MultiNLI | paper SNLI, paper MultiNLI | 277,230 |
| Stack Exchange Duplicate questions (bodies) | 250,519 | |
| Stack Exchange Duplicate questions (titles+bodies) | 250,460 | |
| Sentence Compression | paper | 180,000 |
| Wikihow | paper | 128,542 |
| Altlex | paper | 112,696 |
| Quora Question Triplets | - | 103,663 |
| Simple Wikipedia | paper | 102,225 |
| Natural Questions (NQ) | paper | 100,231 |
| SQuAD2.0 | paper | 87,599 |
| TriviaQA | - | 73,346 |
| Total | 1,170,060,424 |