Setting the file. One moment.
Train Sentence Transformer Make Multilingual Example · Train Sentence Transformers · huggingface/skills · Skills Docs
ContentsBack to the top of the page Losses Cross Encoder
Script Train Sentence Transformer Example
scripts/ train_sentence_transformer_make_multilingual_example.py
Python · 308 lines · 12 KB
16 Cross-lingual retrieval works out of the box because translations sit near
17 each other in the joint space.
18
19 Use when you have a strong English bi-encoder and want a multilingual version
20 but lack in-language supervised data. If you DO have in-language labeled data,
21 train directly with MNRL / CoSENTLoss on it. That usually wins on in-language
22 tasks.
23
24 Data: parallel `(english, non_english)` pairs. `sentence-transformers/parallel-sentences-*`
25 covers many corpora (talks, europarl, tatoeba, wikimatrix, opensubtitles, jw300,
26 news-commentary, ...) with `{src}-{tgt}` subsets. ~500k pairs per language is
27 plenty.
28
29 Picks:
30 - Teacher: any English bi-encoder you want a multilingual copy of
31 (all-mpnet-base-v2, all-MiniLM-L6-v2, BAAI/bge-base-en-v1.5, intfloat/e5-base-v2).
32 - Student: must be multilingual (xlm-roberta-base, paraphrase-multilingual-MiniLM-L12-v2,
33 microsoft/mdeberta-v3-base).
34 - Student dim must match teacher dim, otherwise add a PCA-init Dense projection
35 (see train_sentence_transformer_distillation_example.py).
36 """
37
38 from __future__ import annotations
39
40 import argparse
41 import logging
42 import os
43 from contextlib import nullcontext
44
45 import numpy as np
46 import torch
47 from datasets import DatasetDict, load_dataset
48
49 from sentence_transformers import (
50 SentenceTransformer,
51 SentenceTransformerModelCardData,
52 SentenceTransformerTrainer,
53 SentenceTransformerTrainingArguments,
54 )
55 from sentence_transformers.sentence_transformer.evaluation import (
56 MSEEvaluator,
57 SequentialEvaluator,
58 TranslationEvaluator,
59 )
60 from sentence_transformers.sentence_transformer.losses import MSELoss
61 from sentence_transformers.sentence_transformer.modules import Normalize
62
63
64 def autocast_ctx ():
65 """bf16/fp16 autocast for evaluator calls outside the trainer (which has its own autocast)."""
66 if not torch.cuda.is_available():
67 return nullcontext()
68 dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
69 return torch.autocast( "cuda" , dtype = dtype)
70
71
72 def log_trackio_dashboard ():
73 """Surface the Trackio dashboard URL so the user can watch training live."""
74 try :
75 from huggingface_hub import whoami
76
77 hf_user = whoami().get( "name" )
78 if hf_user:
79 logging.info(
80 f "Trackio dashboard (live training progress): https://huggingface.co/spaces/ { hf_user } /trackio"
81 )
82 except Exception :
83 pass
84
85
86 TEACHER_MODEL_NAME = "sentence-transformers/all-mpnet-base-v2"
87 STUDENT_MODEL_NAME = "FacebookAI/xlm-roberta-base"
88
89 PARALLEL_DATASET = "sentence-transformers/parallel-sentences-talks"
90 SOURCE_LANGUAGE = "en"
91 TARGET_LANGUAGES = ( "de" , "es" , "fr" , "it" )
92
93 MAX_SENTENCES_PER_LANGUAGE = 200_000
94 EVAL_SENTENCES_PER_LANGUAGE = 1_000
95 STUDENT_MAX_SEQ_LENGTH = 128
96
97 OUTPUT_DIR = "models/xlm-roberta-multilingual-from-mpnet"
98 RUN_NAME = "xlm-roberta-multilingual-from-mpnet"
99
100 TEACHER_ENCODE_BATCH_SIZE = 256
101 TRAIN_BATCH_SIZE = 64
102 SMOKE_TEST = os.environ.get( "SMOKE_TEST" ) == "1"
103
104
105 def setup_logging ():
106 """Configure logging + TF32. Tees to logs/{RUN_NAME}.log and silences HTTP spam."""
107 os.makedirs( "logs" , exist_ok = True )
108 logging.basicConfig(
109 format = " %(asctime)s - %(message)s " ,
110 datefmt = "%Y-%m- %d %H:%M:%S" ,
111 level = logging. INFO ,
112 handlers = [logging.StreamHandler(), logging.FileHandler( f "logs/ { RUN_NAME } .log" )],
113 force = True ,
114 )
115 for noisy in ( "httpx" , "httpcore" , "huggingface_hub" , "urllib3" , "filelock" , "fsspec" ):
116 logging.getLogger(noisy).setLevel(logging. WARNING )
117 if torch.cuda.is_available():
118 torch.set_float32_matmul_precision( "high" ) # TF32 on Ampere+, no quality loss
119
120
121 def load_parallel_data () -> tuple[DatasetDict, DatasetDict]:
122 """Load (en, target_lang) parallel pairs for each target language as a DatasetDict."""
123 train_dict = DatasetDict()
124 eval_dict = DatasetDict()
125 for tgt in TARGET_LANGUAGES :
126 subset = f " { SOURCE_LANGUAGE } - { tgt } "
127 try :
128 train_ds = load_dataset( PARALLEL_DATASET , subset, split = "train" )
129 except Exception as exc:
130 logging.error( f "Could not load { PARALLEL_DATASET } / { subset } : { exc } " )
131 continue
132 if len (train_ds) > MAX_SENTENCES_PER_LANGUAGE :
133 train_ds = train_ds.select( range ( MAX_SENTENCES_PER_LANGUAGE ))
134
135 try :
136 eval_ds = load_dataset( PARALLEL_DATASET , subset, split = "dev" ).select( range ( EVAL_SENTENCES_PER_LANGUAGE ))
137 except Exception :
138 split = train_ds.train_test_split( test_size = EVAL_SENTENCES_PER_LANGUAGE , shuffle = True , seed = 12 )
139 train_ds, eval_ds = split[ "train" ], split[ "test" ]
140
141 train_dict[subset] = train_ds
142 eval_dict[subset] = eval_ds
143 if not train_dict:
144 raise SystemExit ( f "No language subsets loaded from { PARALLEL_DATASET } . Check TARGET_LANGUAGES." )
145 return train_dict, eval_dict
146
147
148 def build_evaluator (eval_dict: DatasetDict, teacher: SentenceTransformer) -> SequentialEvaluator:
149 """Per-language MSE + TranslationEvaluator. `main_score_function` averages
150 translation accuracies only. MSE (`negative_mse * 100`) is on a different
151 scale and would break the verdict threshold if mixed in."""
152 sub_evaluators = []
153 for subset, ds in eval_dict.items():
154 sub_evaluators.append(
155 MSEEvaluator(
156 source_sentences = ds[ "english" ],
157 target_sentences = ds[ "non_english" ],
158 name = subset,
159 teacher_model = teacher,
160 batch_size = TEACHER_ENCODE_BATCH_SIZE ,
161 )
162 )
163 sub_evaluators.append(
164 TranslationEvaluator(
165 source_sentences = ds[ "english" ],
166 target_sentences = ds[ "non_english" ],
167 name = subset,
168 batch_size = TEACHER_ENCODE_BATCH_SIZE ,
169 )
170 )
171 # Sub-evaluators alternate MSE / Translation per language: scores[1::2] are the translation accuracies.
172 return SequentialEvaluator(sub_evaluators, main_score_function =lambda scores: float (np.mean(scores[ 1 :: 2 ])))
173
174
175 def main () -> None :
176 parser = argparse.ArgumentParser()
177 parser.add_argument(
178 "--eval-only" , type = str , default = None , help = "Skip training; load this saved model and run only the evaluator."
179 )
180 cli, _ = parser.parse_known_args()
181
182 setup_logging()
183
184 if cli.eval_only:
185 logging.info( f "Eval-only mode: loading model from { cli.eval_only } " )
186 student = SentenceTransformer(cli.eval_only)
187 teacher = SentenceTransformer( TEACHER_MODEL_NAME )
188 _, eval_dict = load_parallel_data()
189 evaluator = build_evaluator(eval_dict, teacher)
190 with autocast_ctx():
191 evaluator(student)
192 return
193
194 logging.info( f "Loading teacher: { TEACHER_MODEL_NAME } " )
195 teacher = SentenceTransformer( TEACHER_MODEL_NAME )
196
197 logging.info( f "Loading student: { STUDENT_MODEL_NAME } " )
198 student = SentenceTransformer(
199 STUDENT_MODEL_NAME ,
200 model_card_data = SentenceTransformerModelCardData(
201 language = [ SOURCE_LANGUAGE , * TARGET_LANGUAGES ],
202 license = "apache-2.0" ,
203 model_name = f " { STUDENT_MODEL_NAME .split( '/' )[ - 1 ] } multilingual from { TEACHER_MODEL_NAME .split( '/' )[ - 1 ] } " ,
204 ),
205 )
206 student.max_seq_length = STUDENT_MAX_SEQ_LENGTH
207 # Match the teacher's final Normalize. MSELoss against unit-norm targets fights student
208 # outputs at norm ~5-10 and can silently regress
209 if any ( isinstance (m, Normalize) for m in teacher) and not any ( isinstance (m, Normalize) for m in student):
210 student.append(Normalize())
211
212 if student.get_embedding_dimension() != teacher.get_embedding_dimension():
213 raise SystemExit (
214 f "Student dim ( { student.get_embedding_dimension() } ) != teacher dim "
215 f "( { teacher.get_embedding_dimension() } ). MSELoss requires matching dims. "
216 "Either pick a student with matching dim, or add a Dense projection layer "
217 "(see train_sentence_transformer_distillation_example.py 'MISMATCHED EMBEDDING DIMS')."
218 )
219
220 logging.info( "Loading parallel data" )
221 train_dict, eval_dict = load_parallel_data()
222 if SMOKE_TEST :
223 logging.info( "SMOKE_TEST=1: trimming each language subset; will run max_steps=1 and skip Hub push" )
224 train_dict = DatasetDict({k: v.select( range ( min ( 50 , len (v)))) for k, v in train_dict.items()})
225 eval_dict = DatasetDict({k: v.select( range ( min ( 20 , len (v)))) for k, v in eval_dict.items()})
226
227 def attach_teacher_label (batch):
228 return {
229 "english" : batch[ "english" ],
230 "non_english" : batch[ "non_english" ],
231 "label" : teacher.encode(batch[ "english" ], batch_size = TEACHER_ENCODE_BATCH_SIZE , show_progress_bar = False ),
232 }
233
234 column_names = list (train_dict.values())[ 0 ].column_names
235 logging.info( "Encoding training English sentences with teacher (cached on disk if you save_to_disk)" )
236 train_dict = train_dict.map(attach_teacher_label, batched = True , batch_size = 10_000 , remove_columns = column_names)
237 eval_dict = eval_dict.map(attach_teacher_label, batched = True , batch_size = 10_000 , remove_columns = column_names)
238
239 loss = MSELoss( model = student)
240
241 evaluator = build_evaluator(eval_dict, teacher)
242 logging.info( "Student baseline (before training):" )
243 with autocast_ctx():
244 baseline_eval = evaluator(student)[ "sequential_score" ]
245
246 args = SentenceTransformerTrainingArguments(
247 output_dir = OUTPUT_DIR ,
248 num_train_epochs = 3 ,
249 max_steps = 1 if SMOKE_TEST else - 1 ,
250 per_device_train_batch_size = TRAIN_BATCH_SIZE ,
251 per_device_eval_batch_size = TRAIN_BATCH_SIZE ,
252 learning_rate = 2e-5 ,
253 weight_decay = 0.01 ,
254 warmup_steps = 0.1 ,
255 bf16 = True ,
256 eval_strategy = "steps" ,
257 eval_steps = 0.1 ,
258 save_strategy = "steps" ,
259 save_steps = 0.1 ,
260 save_total_limit = 2 ,
261 logging_steps = 0.01 ,
262 logging_first_step = True ,
263 load_best_model_at_end = True ,
264 metric_for_best_model = "eval_sequential_score" ,
265 greater_is_better = True ,
266 report_to = "none" if SMOKE_TEST else "trackio" ,
267 run_name = RUN_NAME ,
268 seed = 12 ,
269 )
270
271 trainer = SentenceTransformerTrainer(
272 model = student,
273 args = args,
274 train_dataset = train_dict,
275 eval_dataset = eval_dict,
276 loss = loss,
277 evaluator = evaluator,
278 )
279 if not SMOKE_TEST :
280 log_trackio_dashboard()
281 trainer.train()
282
283 logging.info( "Final student evaluation:" )
284 with autocast_ctx():
285 score = evaluator(student)[ "sequential_score" ]
286 delta = score - baseline_eval
287 verdict = "WIN" if delta >= 0.005 else "MARGINAL" if delta >= 0 else "REGRESSION"
288 logging.info( f "VERDICT: { verdict } | score= { score :.4f} | baseline= { baseline_eval :.4f} | delta= { delta :+.4f} " )
289
290 final_dir = f " { OUTPUT_DIR } /final"
291 student.save_pretrained(final_dir)
292 logging.info( f "Saved to { final_dir } " )
293
294 if SMOKE_TEST :
295 logging.info( "SMOKE_TEST=1: skipping Hub push" )
296 return
297
298 try :
299 commit_url = student.push_to_hub( RUN_NAME )
300 logging.info( f "Pushed model to { commit_url.rsplit( '/commit/' , 1 )[ 0 ] } " )
301 except Exception :
302 import traceback
303
304 logging.error( f "Hub push failed: \n{ traceback.format_exc() } " )
305
306
307 if __name__ == "__main__" :
308 main()