Views
No views yet

Pre-built tokie tokenizer included (tokenizer.tkz). 5x faster tokenization, drop-in replacement for HuggingFace tokenizers.

pip install model2vecfrom_pretrained method:1from model2vec import StaticModel
2# Load a pretrained Model2Vec model
3model = StaticModel.from_pretrained("minishlab/potion-retrieval-32M")
4# Compute text embeddings
5embeddings = model.encode(["Example sentence"])| Model | MTEB Retrieval Score |
|---|---|
| all-MiniLM-L6-v2 | 42.92 |
| potion-retrieval-32M | 35.06 |
| static-retrieval-mrl-en-v1 | 34.95 |
| potion-base-32M | 32.67 |
1@software{minishlab2024model2vec,
2 author = {Stephan Tulkens and {van Dongen}, Thomas},
3 title = {Model2Vec: Fast State-of-the-Art Static Embeddings},
4 year = {2024},
5 publisher = {Zenodo},
6 doi = {10.5281/zenodo.17270888},
7 url = {https://github.com/MinishLab/model2vec},
8 license = {MIT}
9}1import random
2import logging
3from datasets import load_dataset, Dataset, DatasetDict
4from sentence_transformers import (
5 SentenceTransformer,
6 SentenceTransformerTrainer,
7 SentenceTransformerTrainingArguments,
8 SentenceTransformerModelCardData,
9)
10from sentence_transformers.losses import MatryoshkaLoss, MultipleNegativesRankingLoss
11from sentence_transformers.training_args import BatchSamplers, MultiDatasetBatchSamplers
12from sentence_transformers.evaluation import NanoBEIREvaluator
13from sentence_transformers.models.StaticEmbedding import StaticEmbedding
14import wandb
15
16logging.basicConfig(
17 format="%(asctime)s - %(message)s", datefmt="%Y-%m-%d %H:%M:%S", level=logging.INFO
18)
19random.seed(12)
20
21
22def load_train_eval_datasets(factor: int = 1):
23 """
24 Loads train and eval datasets from disk if available. Otherwise, downloads
25 them from Hugging Face, preprocesses, and saves them to disk. If `factor` is
26 greater than 1, returns a fraction (1/factor) of each dataset subset.
27
28 :param factor: The factor by which the data is reduced. If factor=1, no reduction is performed.
29 :return: (train_dataset: DatasetDict, eval_dataset: DatasetDict)
30 """
31 try:
32 # Try loading from disk
33 train_dataset = DatasetDict.load_from_disk("datasets/train_dataset")
34 eval_dataset = DatasetDict.load_from_disk("datasets/eval_dataset")
35 except FileNotFoundError:
36 print("Prebuilt datasets not found on disk. Building from scratch...")
37
38 print("Loading gooaq dataset...")
39 gooaq_dataset = load_dataset("sentence-transformers/gooaq", split="train")
40 gooaq_dataset_dict = gooaq_dataset.train_test_split(test_size=10_000, seed=12)
41 gooaq_train_dataset: Dataset = gooaq_dataset_dict["train"]
42 gooaq_eval_dataset: Dataset = gooaq_dataset_dict["test"]
43 print("Loaded gooaq dataset.")
44
45 print("Loading msmarco dataset...")
46 msmarco_dataset = load_dataset(
47 "sentence-transformers/msmarco-co-condenser-margin-mse-sym-mnrl-mean-v1",
48 "triplet",
49 split="train"
50 )
51 msmarco_dataset_dict = msmarco_dataset.train_test_split(test_size=10_000, seed=12)
52 msmarco_train_dataset: Dataset = msmarco_dataset_dict["train"]
53 msmarco_eval_dataset: Dataset = msmarco_dataset_dict["test"]
54 print("Loaded msmarco dataset.")
55
56 print("Loading squad dataset...")
57 squad_dataset = load_dataset("sentence-transformers/squad", split="train")
58 squad_dataset_dict = squad_dataset.train_test_split(test_size=10_000, seed=12)
59 squad_train_dataset: Dataset = squad_dataset_dict["train"]
60 squad_eval_dataset: Dataset = squad_dataset_dict["test"]
61 print("Loaded squad dataset.")
62
63 print("Loading s2orc dataset...")
64 s2orc_dataset = load_dataset(
65 "sentence-transformers/s2orc",
66 "title-abstract-pair",
67 split="train[:100000]" # limit to 100k
68 )
69 s2orc_dataset_dict = s2orc_dataset.train_test_split(test_size=10_000, seed=12)
70 s2orc_train_dataset: Dataset = s2orc_dataset_dict["train"]
71 s2orc_eval_dataset: Dataset = s2orc_dataset_dict["test"]
72 print("Loaded s2orc dataset.")
73
74 print("Loading allnli dataset...")
75 allnli_train_dataset = load_dataset(
76 "sentence-transformers/all-nli",
77 "triplet",
78 split="train"
79 )
80 allnli_eval_dataset = load_dataset(
81 "sentence-transformers/all-nli",
82 "triplet",
83 split="dev"
84 )
85 print("Loaded allnli dataset.")
86
87 print("Loading paq dataset...")
88 paq_dataset = load_dataset("sentence-transformers/paq", split="train")
89 paq_dataset_dict = paq_dataset.train_test_split(test_size=10_000, seed=12)
90 paq_train_dataset: Dataset = paq_dataset_dict["train"]
91 paq_eval_dataset: Dataset = paq_dataset_dict["test"]
92 print("Loaded paq dataset.")
93
94 print("Loading trivia_qa dataset...")
95 trivia_qa = load_dataset("sentence-transformers/trivia-qa", split="train")
96 trivia_qa_dataset_dict = trivia_qa.train_test_split(test_size=5_000, seed=12)
97 trivia_qa_train_dataset: Dataset = trivia_qa_dataset_dict["train"]
98 trivia_qa_eval_dataset: Dataset = trivia_qa_dataset_dict["test"]
99 print("Loaded trivia_qa dataset.")
100
101 print("Loading msmarco_10m dataset...")
102 msmarco_10m_dataset = load_dataset("bclavie/msmarco-10m-triplets", split="train")
103 msmarco_10m_dataset_dict = msmarco_10m_dataset.train_test_split(
104 test_size=10_000, seed=12
105 )
106 msmarco_10m_train_dataset: Dataset = msmarco_10m_dataset_dict["train"]
107 msmarco_10m_eval_dataset: Dataset = msmarco_10m_dataset_dict["test"]
108 print("Loaded msmarco_10m dataset.")
109
110 print("Loading swim_ir dataset...")
111 swim_ir_dataset = load_dataset(
112 "nthakur/swim-ir-monolingual",
113 "en",
114 split="train"
115 ).select_columns(["query", "text"])
116 swim_ir_dataset_dict = swim_ir_dataset.train_test_split(
117 test_size=10_000, seed=12
118 )
119 swim_ir_train_dataset: Dataset = swim_ir_dataset_dict["train"]
120 swim_ir_eval_dataset: Dataset = swim_ir_dataset_dict["test"]
121 print("Loaded swim_ir dataset.")
122
123 # NOTE: 20 negatives
124 print("Loading pubmedqa dataset...")
125 pubmedqa_dataset = load_dataset(
126 "sentence-transformers/pubmedqa",
127 "triplet-20",
128 split="train"
129 )
130 pubmedqa_dataset_dict = pubmedqa_dataset.train_test_split(test_size=100, seed=12)
131 pubmedqa_train_dataset: Dataset = pubmedqa_dataset_dict["train"]
132 pubmedqa_eval_dataset: Dataset = pubmedqa_dataset_dict["test"]
133 print("Loaded pubmedqa dataset.")
134
135 # NOTE: A lot of overlap with anchor/positives
136 print("Loading miracl dataset...")
137 miracl_dataset = load_dataset(
138 "sentence-transformers/miracl",
139 "en-triplet-all",
140 split="train"
141 )
142 miracl_dataset_dict = miracl_dataset.train_test_split(test_size=10_000, seed=12)
143 miracl_train_dataset: Dataset = miracl_dataset_dict["train"]
144 miracl_eval_dataset: Dataset = miracl_dataset_dict["test"]
145 print("Loaded miracl dataset.")
146
147 # NOTE: A lot of overlap with anchor/positives
148 print("Loading mldr dataset...")
149 mldr_dataset = load_dataset(
150 "sentence-transformers/mldr",
151 "en-triplet-all",
152 split="train"
153 )
154 mldr_dataset_dict = mldr_dataset.train_test_split(test_size=10_000, seed=12)
155 mldr_train_dataset: Dataset = mldr_dataset_dict["train"]
156 mldr_eval_dataset: Dataset = mldr_dataset_dict["test"]
157 print("Loaded mldr dataset.")
158
159 # NOTE: A lot of overlap with anchor/positives
160 print("Loading mr_tydi dataset...")
161 mr_tydi_dataset = load_dataset(
162 "sentence-transformers/mr-tydi",
163 "en-triplet-all",
164 split="train"
165 )
166 mr_tydi_dataset_dict = mr_tydi_dataset.train_test_split(test_size=10_000, seed=12)
167 mr_tydi_train_dataset: Dataset = mr_tydi_dataset_dict["train"]
168 mr_tydi_eval_dataset: Dataset = mr_tydi_dataset_dict["test"]
169 print("Loaded mr_tydi dataset.")
170
171 train_dataset = DatasetDict({
172 "gooaq": gooaq_train_dataset,
173 "msmarco": msmarco_train_dataset,
174 "squad": squad_train_dataset,
175 "s2orc": s2orc_train_dataset,
176 "allnli": allnli_train_dataset,
177 "paq": paq_train_dataset,
178 "trivia_qa": trivia_qa_train_dataset,
179 "msmarco_10m": msmarco_10m_train_dataset,
180 "swim_ir": swim_ir_train_dataset,
181 "pubmedqa": pubmedqa_train_dataset,
182 "miracl": miracl_train_dataset,
183 "mldr": mldr_train_dataset,
184 "mr_tydi": mr_tydi_train_dataset,
185 })
186 eval_dataset = DatasetDict({
187 "gooaq": gooaq_eval_dataset,
188 "msmarco": msmarco_eval_dataset,
189 "squad": squad_eval_dataset,
190 "s2orc": s2orc_eval_dataset,
191 "allnli": allnli_eval_dataset,
192 "paq": paq_eval_dataset,
193 "trivia_qa": trivia_qa_eval_dataset,
194 "msmarco_10m": msmarco_10m_eval_dataset,
195 "swim_ir": swim_ir_eval_dataset,
196 "pubmedqa": pubmedqa_eval_dataset,
197 "miracl": miracl_eval_dataset,
198 "mldr": mldr_eval_dataset,
199 "mr_tydi": mr_tydi_eval_dataset,
200 })
201
202 # Save to disk for next time
203 train_dataset.save_to_disk("datasets/train_dataset")
204 eval_dataset.save_to_disk("datasets/eval_dataset")
205
206 # Quit to avoid memory overhead on large datasets
207 quit()
208
209 # Reduce the dataset if factor > 1
210 if factor > 1:
211 for subset_name in train_dataset:
212 ds = train_dataset[subset_name].shuffle(seed=42)
213 new_len = len(ds) // factor
214 train_dataset[subset_name] = ds.select(range(new_len))
215
216 for subset_name in eval_dataset:
217 ds = eval_dataset[subset_name].shuffle(seed=42)
218 new_len = len(ds) // factor
219 eval_dataset[subset_name] = ds.select(range(new_len))
220
221 return train_dataset, eval_dataset
222
223
224def main():
225 wandb.init(entity="minishlab", project="minishlab")
226 # 1. Load a model to finetune
227 static_embedding = StaticEmbedding.from_model2vec("minishlab/potion-base-32M")
228
229 # 2. Initialize the SentenceTransformer model
230 model_name = "potion-retrieval-32M"
231 model = SentenceTransformer(
232 modules=[static_embedding],
233 model_card_data=SentenceTransformerModelCardData(
234 language="en",
235 license="MIT",
236 model_name=model_name,
237 ),
238 )
239
240 # 3. Load training & evaluation datasets
241 # NOTE: we reduce the total dataset size by a factor of 10
242 train_dataset, eval_dataset = load_train_eval_datasets(factor=10)
243 print(train_dataset)
244
245 # 4. Define a loss function
246 loss = MultipleNegativesRankingLoss(model)
247 loss = MatryoshkaLoss(model, loss, matryoshka_dims=[32, 64, 128, 256, 512])
248
249 # 5. Specify training arguments
250 run_name = model_name
251 epochs = 3
252 lr = 0.05
253 args = SentenceTransformerTrainingArguments(
254 output_dir=f"models/{run_name}",
255 num_train_epochs=epochs,
256 per_device_train_batch_size=2048,
257 per_device_eval_batch_size=2048,
258 learning_rate=lr,
259 warmup_ratio=0.1,
260 fp16=False,
261 bf16=True,
262 batch_sampler=BatchSamplers.NO_DUPLICATES,
263 multi_dataset_batch_sampler=MultiDatasetBatchSamplers.PROPORTIONAL,
264 eval_strategy="steps",
265 eval_steps=250,
266 save_strategy="steps",
267 save_steps=250,
268 save_total_limit=2,
269 logging_steps=250,
270 logging_first_step=True,
271 run_name=run_name,
272 report_to=["wandb"],
273 load_best_model_at_end=True,
274 metric_for_best_model="eval_NanoBEIR_mean_cosine_ndcg@10",
275 greater_is_better=True,
276 )
277
278 # 6. Create an evaluator & evaluate the base model
279 evaluator = NanoBEIREvaluator()
280 evaluator(model)
281
282 # 7. Create a trainer & train
283 trainer = SentenceTransformerTrainer(
284 model=model,
285 args=args,
286 train_dataset=train_dataset,
287 eval_dataset=eval_dataset,
288 loss=loss,
289 evaluator=evaluator,
290 )
291 trainer.train()
292
293 # 8. Evaluate the trained model and save
294 evaluator(model)
295 model.save_pretrained(f"models/{run_name}/final")
296
297
298if __name__ == "__main__":
299 main()