Setting the file. One moment.
Train Sentence Transformer Matryoshka Example · Train Sentence Transformers · huggingface/skills · Skills Docs
ContentsBack to the top of the page Losses Cross Encoder
scripts/ train_sentence_transformer_matryoshka_example.py
Python · 212 lines · 7 KB
16
17 Typical use: train at [768, 512, 256, 128, 64], deploy at 128 for 6x smaller
18 index + 6x faster ANN with minimal quality loss.
19
20 Run locally:
21 pip install "sentence-transformers[train]>=5.0"
22 python train_sentence_transformer_matryoshka_example.py
23 """
24
25 from __future__ import annotations
26
27 import argparse
28 import logging
29 import os
30 from contextlib import nullcontext
31
32 import torch
33 from datasets import load_dataset
34
35 from sentence_transformers import (
36 SentenceTransformer,
37 SentenceTransformerTrainer,
38 SentenceTransformerTrainingArguments,
39 )
40 from sentence_transformers.base.sampler import BatchSamplers
41 from sentence_transformers.sentence_transformer.evaluation import (
42 EmbeddingSimilarityEvaluator,
43 NanoBEIREvaluator,
44 SequentialEvaluator,
45 )
46 from sentence_transformers.sentence_transformer.losses import MatryoshkaLoss, MultipleNegativesRankingLoss
47 from sentence_transformers.util.similarity import SimilarityFunction
48
49
50 def autocast_ctx ():
51 """bf16/fp16 autocast for evaluator calls outside the trainer (which has its own autocast)."""
52 if not torch.cuda.is_available():
53 return nullcontext()
54 dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
55 return torch.autocast( "cuda" , dtype = dtype)
56
57
58 def log_trackio_dashboard ():
59 """Surface the Trackio dashboard URL so the user can watch training live."""
60 try :
61 from huggingface_hub import whoami
62
63 hf_user = whoami().get( "name" )
64 if hf_user:
65 logging.info(
66 f "Trackio dashboard (live training progress): https://huggingface.co/spaces/ { hf_user } /trackio"
67 )
68 except Exception :
69 pass
70
71
72 MODEL_NAME = "microsoft/mpnet-base"
73 MATRYOSHKA_DIMS = [ 768 , 512 , 256 , 128 , 64 ]
74 OUTPUT_DIR = "models/mpnet-matryoshka"
75 RUN_NAME = "mpnet-matryoshka"
76 SMOKE_TEST = os.environ.get( "SMOKE_TEST" ) == "1"
77
78
79 def setup_logging ():
80 """Configure logging + TF32. Tees to logs/{RUN_NAME}.log and silences HTTP spam."""
81 os.makedirs( "logs" , exist_ok = True )
82 logging.basicConfig(
83 format = " %(asctime)s - %(message)s " ,
84 datefmt = "%Y-%m- %d %H:%M:%S" ,
85 level = logging. INFO ,
86 handlers = [logging.StreamHandler(), logging.FileHandler( f "logs/ { RUN_NAME } .log" )],
87 force = True ,
88 )
89 for noisy in ( "httpx" , "httpcore" , "huggingface_hub" , "urllib3" , "filelock" , "fsspec" ):
90 logging.getLogger(noisy).setLevel(logging. WARNING )
91 if torch.cuda.is_available():
92 torch.set_float32_matmul_precision( "high" ) # TF32 on Ampere+, no quality loss
93
94
95 def main () -> None :
96 parser = argparse.ArgumentParser()
97 parser.add_argument(
98 "--eval-only" , type = str , default = None , help = "Skip training; load this saved model and run only the evaluator."
99 )
100 cli, _ = parser.parse_known_args()
101
102 setup_logging()
103
104 if cli.eval_only:
105 logging.info( f "Eval-only mode: loading model from { cli.eval_only } " )
106 model = SentenceTransformer(cli.eval_only)
107 evaluator = NanoBEIREvaluator()
108 with autocast_ctx():
109 evaluator(model)
110 return
111
112 model = SentenceTransformer( MODEL_NAME )
113
114 train_size = 50 if SMOKE_TEST else 50_000
115 eval_size = 20 if SMOKE_TEST else 1_000
116 if SMOKE_TEST :
117 logging.info( "SMOKE_TEST=1: trimmed dataset; will run max_steps=1 and skip Hub push" )
118 train_dataset = load_dataset( "sentence-transformers/all-nli" , "triplet" , split = "train" ).select( range (train_size))
119 eval_dataset = load_dataset( "sentence-transformers/all-nli" , "triplet" , split = "dev" ).select( range (eval_size))
120
121 inner_loss = MultipleNegativesRankingLoss(model)
122 loss = MatryoshkaLoss(model, inner_loss, matryoshka_dims = MATRYOSHKA_DIMS )
123
124 stsb = load_dataset( "sentence-transformers/stsb" , split = "validation" )
125 per_dim_evaluators = [
126 EmbeddingSimilarityEvaluator(
127 sentences1 = stsb[ "sentence1" ],
128 sentences2 = stsb[ "sentence2" ],
129 scores = stsb[ "score" ],
130 main_similarity = SimilarityFunction. COSINE ,
131 name = f "sts-dev- { dim } " ,
132 truncate_dim = dim,
133 )
134 for dim in MATRYOSHKA_DIMS
135 ]
136 evaluator = SequentialEvaluator(
137 [ * per_dim_evaluators, NanoBEIREvaluator()],
138 main_score_function =lambda scores: scores[ 0 ],
139 )
140 logging.info( "Baseline evaluation (before training):" )
141 with autocast_ctx():
142 # Must run before deriving metric_key: each sub-evaluator mutates its primary_metric to add the name_ prefix.
143 baseline_result = evaluator(model)
144 # Drive on the first per-dim evaluator's metric (matches main_score_function above).
145 metric_key = f "eval_ { per_dim_evaluators[ 0 ].primary_metric } "
146 baseline_eval = baseline_result[per_dim_evaluators[ 0 ].primary_metric]
147
148 args = SentenceTransformerTrainingArguments(
149 output_dir = OUTPUT_DIR ,
150 num_train_epochs = 1 ,
151 max_steps = 1 if SMOKE_TEST else - 1 ,
152 per_device_train_batch_size = 128 ,
153 per_device_eval_batch_size = 128 ,
154 learning_rate = 2e-5 ,
155 weight_decay = 0.01 ,
156 warmup_steps = 0.1 ,
157 bf16 = True ,
158 batch_sampler = BatchSamplers. NO_DUPLICATES ,
159 eval_strategy = "steps" ,
160 eval_steps = 0.1 ,
161 save_strategy = "steps" ,
162 save_steps = 0.1 ,
163 save_total_limit = 2 ,
164 logging_steps = 0.01 ,
165 logging_first_step = True ,
166 load_best_model_at_end = True ,
167 metric_for_best_model = metric_key,
168 greater_is_better = True ,
169 report_to = "none" if SMOKE_TEST else "trackio" ,
170 run_name = RUN_NAME ,
171 seed = 12 ,
172 )
173
174 trainer = SentenceTransformerTrainer(
175 model = model,
176 args = args,
177 train_dataset = train_dataset,
178 eval_dataset = eval_dataset,
179 loss = loss,
180 evaluator = evaluator,
181 )
182 if not SMOKE_TEST :
183 log_trackio_dashboard()
184 trainer.train()
185
186 logging.info( "Post-training evaluation:" )
187 with autocast_ctx():
188 score = evaluator(model)[per_dim_evaluators[ 0 ].primary_metric]
189 delta = score - baseline_eval
190 verdict = "WIN" if delta >= 0.005 else "MARGINAL" if delta >= 0 else "REGRESSION"
191 logging.info( f "VERDICT: { verdict } | score= { score :.4f} | baseline= { baseline_eval :.4f} | delta= { delta :+.4f} " )
192
193 final_dir = f " { OUTPUT_DIR } /final"
194 model.save_pretrained(final_dir)
195 logging.info( f "Saved to { final_dir } " )
196 logging.info( f "To use at a specific dimension, load with: SentenceTransformer( { final_dir !r} , truncate_dim=128)" )
197
198 if SMOKE_TEST :
199 logging.info( "SMOKE_TEST=1: skipping Hub push" )
200 return
201
202 try :
203 commit_url = model.push_to_hub( RUN_NAME )
204 logging.info( f "Pushed model to { commit_url.rsplit( '/commit/' , 1 )[ 0 ] } " )
205 except Exception :
206 import traceback
207
208 logging.error( f "Hub push failed: \n{ traceback.format_exc() } " )
209
210
211 if __name__ == "__main__" :
212 main()