README.md
9.9 KB · 171 lines · markdown Raw
1 ---
2 language: en
3 license: apache-2.0
4 library_name: sentence-transformers
5 tags:
6 - sentence-transformers
7 - feature-extraction
8 - sentence-similarity
9 - transformers
10 datasets:
11 - s2orc
12 - flax-sentence-embeddings/stackexchange_xml
13 - ms_marco
14 - gooaq
15 - yahoo_answers_topics
16 - code_search_net
17 - search_qa
18 - eli5
19 - snli
20 - multi_nli
21 - wikihow
22 - natural_questions
23 - trivia_qa
24 - embedding-data/sentence-compression
25 - embedding-data/flickr30k-captions
26 - embedding-data/altlex
27 - embedding-data/simple-wiki
28 - embedding-data/QQP
29 - embedding-data/SPECTER
30 - embedding-data/PAQ_pairs
31 - embedding-data/WikiAnswers
32 pipeline_tag: sentence-similarity
33 ---
34
35
36 # all-distilroberta-v1
37 This is a [sentence-transformers](https://www.SBERT.net) model: It maps sentences & paragraphs to a 768 dimensional dense vector space and can be used for tasks like clustering or semantic search.
38
39 ## Usage (Sentence-Transformers)
40 Using this model becomes easy when you have [sentence-transformers](https://www.SBERT.net) installed:
41
42 ```
43 pip install -U sentence-transformers
44 ```
45
46 Then you can use the model like this:
47 ```python
48 from sentence_transformers import SentenceTransformer
49 sentences = ["This is an example sentence", "Each sentence is converted"]
50
51 model = SentenceTransformer('sentence-transformers/all-distilroberta-v1')
52 embeddings = model.encode(sentences)
53 print(embeddings)
54 ```
55
56 ## Usage (HuggingFace Transformers)
57 Without [sentence-transformers](https://www.SBERT.net), you can use the model like this: First, you pass your input through the transformer model, then you have to apply the right pooling-operation on-top of the contextualized word embeddings.
58
59 ```python
60 from transformers import AutoTokenizer, AutoModel
61 import torch
62 import torch.nn.functional as F
63
64 #Mean Pooling - Take attention mask into account for correct averaging
65 def mean_pooling(model_output, attention_mask):
66 token_embeddings = model_output[0] #First element of model_output contains all token embeddings
67 input_mask_expanded = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
68 return torch.sum(token_embeddings * input_mask_expanded, 1) / torch.clamp(input_mask_expanded.sum(1), min=1e-9)
69
70
71 # Sentences we want sentence embeddings for
72 sentences = ['This is an example sentence', 'Each sentence is converted']
73
74 # Load model from HuggingFace Hub
75 tokenizer = AutoTokenizer.from_pretrained('sentence-transformers/all-distilroberta-v1')
76 model = AutoModel.from_pretrained('sentence-transformers/all-distilroberta-v1')
77
78 # Tokenize sentences
79 encoded_input = tokenizer(sentences, padding=True, truncation=True, return_tensors='pt')
80
81 # Compute token embeddings
82 with torch.no_grad():
83 model_output = model(**encoded_input)
84
85 # Perform pooling
86 sentence_embeddings = mean_pooling(model_output, encoded_input['attention_mask'])
87
88 # Normalize embeddings
89 sentence_embeddings = F.normalize(sentence_embeddings, p=2, dim=1)
90
91 print("Sentence embeddings:")
92 print(sentence_embeddings)
93 ```
94
95 ------
96
97 ## Background
98
99 The project aims to train sentence embedding models on very large sentence level datasets using a self-supervised
100 contrastive learning objective. We used the pretrained [`distilroberta-base`](https://huggingface.co/distilroberta-base) model and fine-tuned in on a
101 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.
102
103 We developped this model during the
104 [Community week using JAX/Flax for NLP & CV](https://discuss.huggingface.co/t/open-to-the-community-community-week-using-jax-flax-for-nlp-cv/7104),
105 organized by Hugging Face. We developped this model as part of the project:
106 [Train the Best Sentence Embedding Model Ever with 1B Training Pairs](https://discuss.huggingface.co/t/train-the-best-sentence-embedding-model-ever-with-1b-training-pairs/7354). We benefited from efficient hardware infrastructure to run the project: 7 TPUs v3-8, as well as intervention from Googles Flax, JAX, and Cloud team member about efficient deep learning frameworks.
107
108 ## Intended uses
109
110 Our model is intented to be used as a sentence and short paragraph encoder. Given an input text, it ouptuts a vector which captures
111 the semantic information. The sentence vector may be used for information retrieval, clustering or sentence similarity tasks.
112
113 By default, input text longer than 128 word pieces is truncated.
114
115
116 ## Training procedure
117
118 ### Pre-training
119
120 We use the pretrained [`distilroberta-base`](https://huggingface.co/distilroberta-base). Please refer to the model card for more detailed information about the pre-training procedure.
121
122 ### Fine-tuning
123
124 We fine-tune the model using a contrastive objective. Formally, we compute the cosine similarity from each possible sentence pairs from the batch.
125 We then apply the cross entropy loss by comparing with true pairs.
126
127 #### Hyper parameters
128
129 We trained ou model on a TPU v3-8. We train the model during 920k steps using a batch size of 512 (64 per TPU core).
130 We use a learning rate warm up of 500. The sequence length was limited to 128 tokens. We used the AdamW optimizer with
131 a 2e-5 learning rate. The full training script is accessible in this current repository: `train_script.py`.
132
133 #### Training data
134
135 We use the concatenation from multiple datasets to fine-tune our model. The total number of sentence pairs is above 1 billion sentences.
136 We sampled each dataset given a weighted probability which configuration is detailed in the `data_config.json` file.
137
138
139 | Dataset | Paper | Number of training tuples |
140 |--------------------------------------------------------|:----------------------------------------:|:--------------------------:|
141 | [Reddit comments (2015-2018)](https://github.com/PolyAI-LDN/conversational-datasets/tree/master/reddit) | [paper](https://arxiv.org/abs/1904.06472) | 726,484,430 |
142 | [S2ORC](https://github.com/allenai/s2orc) Citation pairs (Abstracts) | [paper](https://aclanthology.org/2020.acl-main.447/) | 116,288,806 |
143 | [WikiAnswers](https://github.com/afader/oqa#wikianswers-corpus) Duplicate question pairs | [paper](https://doi.org/10.1145/2623330.2623677) | 77,427,422 |
144 | [PAQ](https://github.com/facebookresearch/PAQ) (Question, Answer) pairs | [paper](https://arxiv.org/abs/2102.07033) | 64,371,441 |
145 | [S2ORC](https://github.com/allenai/s2orc) Citation pairs (Titles) | [paper](https://aclanthology.org/2020.acl-main.447/) | 52,603,982 |
146 | [S2ORC](https://github.com/allenai/s2orc) (Title, Abstract) | [paper](https://aclanthology.org/2020.acl-main.447/) | 41,769,185 |
147 | [Stack Exchange](https://huggingface.co/datasets/flax-sentence-embeddings/stackexchange_xml) (Title, Body) pairs | - | 25,316,456 |
148 | [MS MARCO](https://microsoft.github.io/msmarco/) triplets | [paper](https://doi.org/10.1145/3404835.3462804) | 9,144,553 |
149 | [GOOAQ: Open Question Answering with Diverse Answer Types](https://github.com/allenai/gooaq) | [paper](https://arxiv.org/pdf/2104.08727.pdf) | 3,012,496 |
150 | [Yahoo Answers](https://www.kaggle.com/soumikrakshit/yahoo-answers-dataset) (Title, Answer) | [paper](https://proceedings.neurips.cc/paper/2015/hash/250cf8b51c773f3f8dc8b4be867a9a02-Abstract.html) | 1,198,260 |
151 | [Code Search](https://huggingface.co/datasets/code_search_net) | - | 1,151,414 |
152 | [COCO](https://cocodataset.org/#home) Image captions | [paper](https://link.springer.com/chapter/10.1007%2F978-3-319-10602-1_48) | 828,395|
153 | [SPECTER](https://github.com/allenai/specter) citation triplets | [paper](https://doi.org/10.18653/v1/2020.acl-main.207) | 684,100 |
154 | [Yahoo Answers](https://www.kaggle.com/soumikrakshit/yahoo-answers-dataset) (Question, Answer) | [paper](https://proceedings.neurips.cc/paper/2015/hash/250cf8b51c773f3f8dc8b4be867a9a02-Abstract.html) | 681,164 |
155 | [Yahoo Answers](https://www.kaggle.com/soumikrakshit/yahoo-answers-dataset) (Title, Question) | [paper](https://proceedings.neurips.cc/paper/2015/hash/250cf8b51c773f3f8dc8b4be867a9a02-Abstract.html) | 659,896 |
156 | [SearchQA](https://huggingface.co/datasets/search_qa) | [paper](https://arxiv.org/abs/1704.05179) | 582,261 |
157 | [Eli5](https://huggingface.co/datasets/eli5) | [paper](https://doi.org/10.18653/v1/p19-1346) | 325,475 |
158 | [Flickr 30k](https://shannon.cs.illinois.edu/DenotationGraph/) | [paper](https://transacl.org/ojs/index.php/tacl/article/view/229/33) | 317,695 |
159 | [Stack Exchange](https://huggingface.co/datasets/flax-sentence-embeddings/stackexchange_xml) Duplicate questions (titles) | | 304,525 |
160 | AllNLI ([SNLI](https://nlp.stanford.edu/projects/snli/) and [MultiNLI](https://cims.nyu.edu/~sbowman/multinli/) | [paper SNLI](https://doi.org/10.18653/v1/d15-1075), [paper MultiNLI](https://doi.org/10.18653/v1/n18-1101) | 277,230 |
161 | [Stack Exchange](https://huggingface.co/datasets/flax-sentence-embeddings/stackexchange_xml) Duplicate questions (bodies) | | 250,519 |
162 | [Stack Exchange](https://huggingface.co/datasets/flax-sentence-embeddings/stackexchange_xml) Duplicate questions (titles+bodies) | | 250,460 |
163 | [Sentence Compression](https://github.com/google-research-datasets/sentence-compression) | [paper](https://www.aclweb.org/anthology/D13-1155/) | 180,000 |
164 | [Wikihow](https://github.com/pvl/wikihow_pairs_dataset) | [paper](https://arxiv.org/abs/1810.09305) | 128,542 |
165 | [Altlex](https://github.com/chridey/altlex/) | [paper](https://aclanthology.org/P16-1135.pdf) | 112,696 |
166 | [Quora Question Triplets](https://quoradata.quora.com/First-Quora-Dataset-Release-Question-Pairs) | - | 103,663 |
167 | [Simple Wikipedia](https://cs.pomona.edu/~dkauchak/simplification/) | [paper](https://www.aclweb.org/anthology/P11-2117/) | 102,225 |
168 | [Natural Questions (NQ)](https://ai.google.com/research/NaturalQuestions) | [paper](https://transacl.org/ojs/index.php/tacl/article/view/1455) | 100,231 |
169 | [SQuAD2.0](https://rajpurkar.github.io/SQuAD-explorer/) | [paper](https://aclanthology.org/P18-2124.pdf) | 87,599 |
170 | [TriviaQA](https://huggingface.co/datasets/trivia_qa) | - | 73,346 |
171 | **Total** | | **1,124,818,467** |