Setting the file. One moment. Train Sparse Encoder Example · Train Sentence Transformers · huggingface/skills · Skills DocsLosses Cross Encoder
scripts/train_sparse_encoder_example.py
scripts/train_sparse_encoder_example.py
Python·237 lines·8 KB
- SparseNanoBEIREvaluator for sparse retrieval metrics
17- load_best_model_at_end on the retrieval metric
18
19Base model must expose a masked-LM head. Any `AutoModelForMaskedLM`-compatible
20checkpoint works (DistilBERT, BERT, MiniLM MLM variants, existing SPLADE models).
21
22Run locally:
23 pip install "sentence-transformers[train]>=5.0"
24 python train_sparse_encoder_example.py
25
26Multi-GPU:
27 accelerate launch train_sparse_encoder_example.py
28
29Hugging Face Jobs: paste this file's contents as the `script` in hf_jobs(...).
30"""
31
32from __future__ import annotations
33
34import argparse
35import logging
36import os
37from contextlib import nullcontext
38
39import torch
40from datasets import load_dataset
41
42from sentence_transformers import (
43 SparseEncoder,
44 SparseEncoderModelCardData,
45 SparseEncoderTrainer,
46 SparseEncoderTrainingArguments,
47)
48from sentence_transformers.base.sampler import BatchSamplers
49from sentence_transformers.sparse_encoder.evaluation import SparseNanoBEIREvaluator
50from sentence_transformers.sparse_encoder.losses import (
51 SparseMultipleNegativesRankingLoss,
52 SpladeLoss,
53)
54
55
56def autocast_ctx():
57 """bf16/fp16 autocast for evaluator calls outside the trainer (which has its own autocast)."""
58 if not torch.cuda.is_available():
59 return nullcontext()
60 dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
61 return torch.autocast("cuda", dtype=dtype)
62
63
64def log_trackio_dashboard():
65 """Surface the Trackio dashboard URL so the user can watch training live."""
66 try:
67 from huggingface_hub import whoami
68
69 hf_user = whoami().get("name")
70 if hf_user:
71 logging.info(
72 f"Trackio dashboard (live training progress): https://huggingface.co/spaces/{hf_user}/trackio"
73 )
74 except Exception:
75 pass
76
77
78MODEL_NAME = "distilbert/distilbert-base-uncased"
79DATASET_NAME = "sentence-transformers/gooaq"
80TRAIN_SIZE = 100_000
81EVAL_SIZE = 1_000
82OUTPUT_DIR = "models/distilbert-splade-gooaq"
83RUN_NAME = "distilbert-splade-gooaq"
84
85QUERY_REGULARIZER_WEIGHT = 5e-5
86DOCUMENT_REGULARIZER_WEIGHT = 3e-5
87SMOKE_TEST = os.environ.get("SMOKE_TEST") == "1"
88
89
90def setup_logging():
91 """Configure logging + TF32. Tees to logs/{RUN_NAME}.log and silences HTTP spam."""
92 os.makedirs("logs", exist_ok=True)
93 logging.basicConfig(
94 format="%(asctime)s - %(message)s",
95 datefmt="%Y-%m-%d %H:%M:%S",
96 level=logging.INFO,
97 handlers=[logging.StreamHandler(), logging.FileHandler(f"logs/{RUN_NAME}.log")],
98 force=True,
99 )
100 for noisy in ("httpx", "httpcore", "huggingface_hub", "urllib3", "filelock", "fsspec"):
101 logging.getLogger(noisy).setLevel(logging.WARNING)
102 if torch.cuda.is_available():
103 torch.set_float32_matmul_precision("high") # TF32 on Ampere+, no quality loss
104
105
106def main() -> None:
107 parser = argparse.ArgumentParser()
108 parser.add_argument(
109 "--eval-only", type=str, default=None, help="Skip training; load this saved model and run only the evaluator."
110 )
111 cli, _ = parser.parse_known_args()
112
113 setup_logging()
114
115 if cli.eval_only:
116 logging.info(f"Eval-only mode: loading model from {cli.eval_only}")
117 model = SparseEncoder(cli.eval_only)
118 evaluator = SparseNanoBEIREvaluator(dataset_names=["msmarco", "nfcorpus", "nq"])
119 with autocast_ctx():
120 evaluator(model)
121 return
122
123 logging.info(f"Loading base model: {MODEL_NAME}")
124 # Prompts are optional for SPLADE: most BERT-style MLM bases don't need them.
125 # If you're starting from a CSR (`Transformer + Pooling + SparseAutoEncoder`)
126 # base like `tomaarsen/csr-mxbai-embed-large-v1-nq` that *was* trained with
127 # prompts, mirror them here to preserve quality:
128 # prompts={"query": "Represent this sentence for similarity: ", "document": ""},
129 # default_prompt_name="document",
130 model = SparseEncoder(
131 MODEL_NAME,
132 model_card_data=SparseEncoderModelCardData(
133 language="en",
134 license="apache-2.0",
135 model_name=f"SPLADE from {MODEL_NAME.split('/')[-1]} trained on GooAQ",
136 ),
137 )
138
139 logging.info(f"Loading dataset: {DATASET_NAME}")
140 train_size = 50 if SMOKE_TEST else TRAIN_SIZE
141 eval_size = 20 if SMOKE_TEST else EVAL_SIZE
142 if SMOKE_TEST:
143 logging.info("SMOKE_TEST=1: trimmed dataset; will run max_steps=1 and skip Hub push")
144 full = load_dataset(DATASET_NAME, split="train")
145 split = full.train_test_split(test_size=eval_size, seed=12)
146 train_dataset = split["train"].select(range(min(train_size, len(split["train"]))))
147 eval_dataset = split["test"]
148 logging.info(f" train: {len(train_dataset):,} rows | eval: {len(eval_dataset):,} rows")
149 logging.info(f" columns: {train_dataset.column_names}")
150
151 loss = SpladeLoss(
152 model=model,
153 loss=SparseMultipleNegativesRankingLoss(model=model),
154 query_regularizer_weight=QUERY_REGULARIZER_WEIGHT,
155 document_regularizer_weight=DOCUMENT_REGULARIZER_WEIGHT,
156 )
157
158 evaluator = SparseNanoBEIREvaluator(dataset_names=["msmarco", "nfcorpus", "nq"])
159 logging.info("Baseline evaluation:")
160 with autocast_ctx():
161 # Must run before deriving metric_key: evaluator(model) mutates primary_metric to add the name_ prefix.
162 baseline_result = evaluator(model)
163 baseline_eval = baseline_result[evaluator.primary_metric]
164 metric_key = f"eval_{evaluator.primary_metric}"
165
166 args = SparseEncoderTrainingArguments(
167 output_dir=OUTPUT_DIR,
168 num_train_epochs=1,
169 max_steps=1 if SMOKE_TEST else -1,
170 per_device_train_batch_size=32,
171 per_device_eval_batch_size=32,
172 learning_rate=2e-5,
173 weight_decay=0.01,
174 warmup_steps=0.1,
175 lr_scheduler_type="linear",
176 bf16=True,
177 batch_sampler=BatchSamplers.NO_DUPLICATES,
178 eval_strategy="steps",
179 eval_steps=0.1,
180 save_strategy="steps",
181 save_steps=0.1,
182 save_total_limit=2,
183 logging_steps=0.01,
184 logging_first_step=True,
185 load_best_model_at_end=True,
186 metric_for_best_model=metric_key,
187 greater_is_better=True,
188 report_to="none" if SMOKE_TEST else "trackio",
189 run_name=RUN_NAME,
190 seed=12,
191 )
192
193 trainer = SparseEncoderTrainer(
194 model=model,
195 args=args,
196 train_dataset=train_dataset,
197 eval_dataset=eval_dataset,
198 loss=loss,
199 evaluator=evaluator,
200 )
201 if not SMOKE_TEST:
202 log_trackio_dashboard()
203 trainer.train()
204
205 logging.info("Post-training evaluation:")
206 with autocast_ctx():
207 result = evaluator(model)
208 score = result[evaluator.primary_metric]
209 delta = score - baseline_eval
210 verdict = "WIN" if delta >= 0.005 else "MARGINAL" if delta >= 0 else "REGRESSION"
211 # Active-dim keys come back name-prefixed (e.g. "NanoBEIR_..._query_active_dims"). Suffix-match for compat.
212 qad = next((v for k, v in result.items() if k.endswith("query_active_dims")), "n/a")
213 cad = next((v for k, v in result.items() if k.endswith("corpus_active_dims")), "n/a")
214 logging.info(
215 f"VERDICT: {verdict} | score={score:.4f} | baseline={baseline_eval:.4f} | delta={delta:+.4f} "
216 f"| query_active={qad} corpus_active={cad}"
217 )
218
219 final_dir = f"{OUTPUT_DIR}/final"
220 model.save_pretrained(final_dir)
221 logging.info(f"Saved final model to {final_dir}")
222
223 if SMOKE_TEST:
224 logging.info("SMOKE_TEST=1: skipping Hub push")
225 return
226
227 try:
228 commit_url = model.push_to_hub(RUN_NAME)
229 logging.info(f"Pushed model to {commit_url.rsplit('/commit/', 1)[0]}")
230 except Exception:
231 import traceback
232
233 logging.error(f"Hub push failed:\n{traceback.format_exc()}")
234
235
236if __name__ == "__main__":
237 main()