Setting the file. One moment. Train Sentence Transformer Example · Train Sentence Transformers · huggingface/skills · Skills DocsLosses Cross Encoder
scripts/train_sentence_transformer_example.py
Python·227 lines·8 KB
- BatchSamplers.NO_DUPLICATES (critical for MNRL)
17- load_best_model_at_end with a retrieval metric
18- Auto model card + optional Hub push
19
20Runs identically in two modes:
21
22 # Local
23 pip install "sentence-transformers[train]>=5.0"
24 python train_sentence_transformer_example.py
25
26 # Or with uv (no explicit install needed)
27 uv run train_sentence_transformer_example.py
28
29 # Multi-GPU
30 accelerate launch train_sentence_transformer_example.py
31
32 # Hugging Face Jobs (paste the entire file contents as `script`)
33 hf_jobs("uv", {
34 "script": "<contents of this file>",
35 "flavor": "a10g-large",
36 "timeout": "3h",
37 "secrets": {"HF_TOKEN": "$HF_TOKEN"},
38 })
39
40Adjust MODEL_NAME, DATASET_NAME, OUTPUT_DIR, RUN_NAME at the top of the script.
41Default Hub push: at end of run, public, under your authenticated user as
42`{user}/{RUN_NAME}`, wrapped in try/except. To skip the push, comment out the
43push_to_hub call. For HF Jobs (ephemeral env), also enable in-trainer push:
44add `push_to_hub=True`, `hub_model_id=RUN_NAME`, `hub_strategy="every_save"`
45to TrainingArguments.
46"""
47
48from __future__ import annotations
49
50import argparse
51import logging
52import os
53from contextlib import nullcontext
54
55import torch
56from datasets import load_dataset
57
58from sentence_transformers import (
59 SentenceTransformer,
60 SentenceTransformerModelCardData,
61 SentenceTransformerTrainer,
62 SentenceTransformerTrainingArguments,
63)
64from sentence_transformers.base.sampler import BatchSamplers
65from sentence_transformers.sentence_transformer.evaluation import NanoBEIREvaluator
66from sentence_transformers.sentence_transformer.losses import MultipleNegativesRankingLoss
67
68
69def autocast_ctx():
70 """bf16/fp16 autocast for evaluator calls outside the trainer (which has its own autocast)."""
71 if not torch.cuda.is_available():
72 return nullcontext()
73 dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
74 return torch.autocast("cuda", dtype=dtype)
75
76
77def log_trackio_dashboard():
78 """Surface the Trackio dashboard URL so the user can watch training live."""
79 try:
80 from huggingface_hub import whoami
81
82 hf_user = whoami().get("name")
83 if hf_user:
84 logging.info(
85 f"Trackio dashboard (live training progress): https://huggingface.co/spaces/{hf_user}/trackio"
86 )
87 except Exception:
88 pass
89
90
91MODEL_NAME = "microsoft/mpnet-base"
92DATASET_NAME = "sentence-transformers/all-nli"
93DATASET_SUBSET = "triplet"
94TRAIN_SIZE = 50_000
95EVAL_SIZE = 1_000
96OUTPUT_DIR = "models/mpnet-base-all-nli"
97RUN_NAME = "mpnet-base-all-nli"
98SMOKE_TEST = os.environ.get("SMOKE_TEST") == "1"
99
100
101def setup_logging():
102 """Configure logging + TF32. Tees to logs/{RUN_NAME}.log and silences HTTP spam."""
103 os.makedirs("logs", exist_ok=True)
104 logging.basicConfig(
105 format="%(asctime)s - %(message)s",
106 datefmt="%Y-%m-%d %H:%M:%S",
107 level=logging.INFO,
108 handlers=[logging.StreamHandler(), logging.FileHandler(f"logs/{RUN_NAME}.log")],
109 force=True,
110 )
111 for noisy in ("httpx", "httpcore", "huggingface_hub", "urllib3", "filelock", "fsspec"):
112 logging.getLogger(noisy).setLevel(logging.WARNING)
113 if torch.cuda.is_available():
114 torch.set_float32_matmul_precision("high") # TF32 on Ampere+, no quality loss
115
116
117def main() -> None:
118 parser = argparse.ArgumentParser()
119 parser.add_argument(
120 "--eval-only", type=str, default=None, help="Skip training; load this saved model and run only the evaluator."
121 )
122 cli, _ = parser.parse_known_args()
123
124 setup_logging()
125
126 if cli.eval_only:
127 logging.info(f"Eval-only mode: loading model from {cli.eval_only}")
128 model = SentenceTransformer(cli.eval_only)
129 evaluator = NanoBEIREvaluator()
130 with autocast_ctx():
131 evaluator(model)
132 return
133
134 logging.info(f"Loading base model: {MODEL_NAME}")
135 model = SentenceTransformer(
136 MODEL_NAME,
137 model_card_data=SentenceTransformerModelCardData(
138 language="en",
139 license="apache-2.0",
140 model_name=f"{MODEL_NAME.split('/')[-1]} finetuned on AllNLI",
141 ),
142 )
143
144 logging.info(f"Loading dataset: {DATASET_NAME} ({DATASET_SUBSET})")
145 train_size = 50 if SMOKE_TEST else TRAIN_SIZE
146 eval_size = 20 if SMOKE_TEST else EVAL_SIZE
147 train_dataset = load_dataset(DATASET_NAME, DATASET_SUBSET, split="train").select(range(train_size))
148 eval_dataset = load_dataset(DATASET_NAME, DATASET_SUBSET, split="dev").select(range(eval_size))
149 if SMOKE_TEST:
150 logging.info("SMOKE_TEST=1: trimmed dataset; will run max_steps=1 and skip Hub push")
151 logging.info(f" train: {len(train_dataset):,} examples")
152 logging.info(f" eval: {len(eval_dataset):,} examples")
153
154 loss = MultipleNegativesRankingLoss(model)
155
156 evaluator = NanoBEIREvaluator()
157 logging.info("Baseline evaluation (before training):")
158 with autocast_ctx():
159 # Must run before deriving metric_key: evaluator(model) mutates primary_metric to add the name_ prefix.
160 baseline_eval = evaluator(model)[evaluator.primary_metric]
161 metric_key = f"eval_{evaluator.primary_metric}"
162
163 args = SentenceTransformerTrainingArguments(
164 output_dir=OUTPUT_DIR,
165 num_train_epochs=1,
166 max_steps=1 if SMOKE_TEST else -1,
167 per_device_train_batch_size=64,
168 per_device_eval_batch_size=64,
169 learning_rate=2e-5,
170 weight_decay=0.01,
171 warmup_steps=0.1,
172 lr_scheduler_type="linear",
173 bf16=True,
174 batch_sampler=BatchSamplers.NO_DUPLICATES,
175 eval_strategy="steps",
176 eval_steps=0.1,
177 save_strategy="steps",
178 save_steps=0.1,
179 save_total_limit=2,
180 logging_steps=0.01,
181 logging_first_step=True,
182 load_best_model_at_end=True,
183 metric_for_best_model=metric_key,
184 greater_is_better=True,
185 report_to="none" if SMOKE_TEST else "trackio",
186 run_name=RUN_NAME,
187 seed=12,
188 )
189
190 trainer = SentenceTransformerTrainer(
191 model=model,
192 args=args,
193 train_dataset=train_dataset,
194 eval_dataset=eval_dataset,
195 loss=loss,
196 evaluator=evaluator,
197 )
198 if not SMOKE_TEST:
199 log_trackio_dashboard()
200 trainer.train()
201
202 logging.info("Post-training evaluation:")
203 with autocast_ctx():
204 score = evaluator(model)[evaluator.primary_metric]
205 delta = score - baseline_eval
206 verdict = "WIN" if delta >= 0.005 else "MARGINAL" if delta >= 0 else "REGRESSION"
207 logging.info(f"VERDICT: {verdict} | score={score:.4f} | baseline={baseline_eval:.4f} | delta={delta:+.4f}")
208
209 final_dir = f"{OUTPUT_DIR}/final"
210 model.save_pretrained(final_dir)
211 logging.info(f"Saved final model to {final_dir}")
212
213 if SMOKE_TEST:
214 logging.info("SMOKE_TEST=1: skipping Hub push")
215 return
216
217 try:
218 commit_url = model.push_to_hub(RUN_NAME) # public by default. Uses your authenticated user
219 logging.info(f"Pushed model to {commit_url.rsplit('/commit/', 1)[0]}")
220 except Exception:
221 import traceback
222
223 logging.error(f"Hub push failed:\n{traceback.format_exc()}")
224
225
226if __name__ == "__main__":
227 main()