Setting the file. One moment. Train Cross Encoder Distillation Example · Train Sentence Transformers · huggingface/skills · Skills DocsLosses Cross Encoder
Next
Script Train Cross Encoder Example
scripts/train_cross_encoder_distillation_example.py
Python·254 lines·9 KB
16
the teacher at training time (precomputed score diffs).
17
18For Embedding-MSE / Listwise-KL variants and the broader pattern map, see the
19sibling `train_sentence_transformer_distillation_example.py` docstring.
20
21CRITICAL: `activation_fn=nn.Identity()` is mandatory. The default `Sigmoid` (with
22`num_labels=1`) saturates raw logits >5 to ~1.0 inside `predict()` at eval time,
23silently collapsing eval ranking (training loss stays healthy while nDCG drops
24from e.g. ~0.59 to ~0.14). See `../references/troubleshooting.md` ("CrossEncoder
25eval nDCG crashes after distillation / listwise / pairwise training").
26
27Data: this script uses `sentence-transformers/msmarco` (`bert-ensemble-margin-mse`
28subset), which has precomputed teacher score diffs per (q, pos, neg) row. To
29distill from your own teacher, replace the dataset loading with a one-time
30teacher pass over your (q, pos, neg) triples and store `score_diff = teacher_pos
31- teacher_neg` as the label.
32
33Run locally:
34 pip install "sentence-transformers[train]>=5.0"
35 python train_cross_encoder_distillation_example.py
36
37Multi-GPU:
38 accelerate launch train_cross_encoder_distillation_example.py
39
40Hugging Face Jobs: paste this file's contents as the `script` in hf_jobs(...).
41"""
42
43from __future__ import annotations
44
45import argparse
46import logging
47import os
48from contextlib import nullcontext
49
50import torch
51import torch.nn as nn
52from datasets import load_dataset, load_from_disk
53from transformers import EarlyStoppingCallback
54
55from sentence_transformers import (
56 CrossEncoder,
57 CrossEncoderModelCardData,
58 CrossEncoderTrainer,
59 CrossEncoderTrainingArguments,
60)
61from sentence_transformers.cross_encoder.evaluation import CrossEncoderNanoBEIREvaluator
62from sentence_transformers.cross_encoder.losses import MarginMSELoss
63
64
65def autocast_ctx():
66 """bf16/fp16 autocast for evaluator calls outside the trainer (which has its own autocast)."""
67 if not torch.cuda.is_available():
68 return nullcontext()
69 dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
70 return torch.autocast("cuda", dtype=dtype)
71
72
73def log_trackio_dashboard():
74 """Surface the Trackio dashboard URL so the user can watch training live."""
75 try:
76 from huggingface_hub import whoami
77
78 hf_user = whoami().get("name")
79 if hf_user:
80 logging.info(
81 f"Trackio dashboard (live training progress): https://huggingface.co/spaces/{hf_user}/trackio"
82 )
83 except Exception:
84 pass
85
86
87MODEL_NAME = "microsoft/MiniLM-L12-H384-uncased"
88DATASET_NAME = "sentence-transformers/msmarco"
89DATASET_SUBSET = "bert-ensemble-margin-mse"
90TRAIN_SIZE = 100_000
91EVAL_SIZE = 5_000
92OUTPUT_DIR = "models/minilm-msmarco-distilled"
93RUN_NAME = "minilm-msmarco-distilled"
94DATA_CACHE = f"data/{RUN_NAME}-resolved"
95SMOKE_TEST = os.environ.get("SMOKE_TEST") == "1"
96
97
98def setup_logging():
99 """Configure logging + TF32. Tees to logs/{RUN_NAME}.log and silences HTTP spam."""
100 os.makedirs("logs", exist_ok=True)
101 logging.basicConfig(
102 format="%(asctime)s - %(message)s",
103 datefmt="%Y-%m-%d %H:%M:%S",
104 level=logging.INFO,
105 handlers=[logging.StreamHandler(), logging.FileHandler(f"logs/{RUN_NAME}.log")],
106 force=True,
107 )
108 for noisy in ("httpx", "httpcore", "huggingface_hub", "urllib3", "filelock", "fsspec"):
109 logging.getLogger(noisy).setLevel(logging.WARNING)
110 if torch.cuda.is_available():
111 torch.set_float32_matmul_precision("high")
112
113
114def load_resolved_dataset():
115 """Load (query, positive, negative, score) rows. The MSMARCO subset is keyed by
116 passage_id / query_id. Resolve to text once and cache to disk so reruns skip the work."""
117 if os.path.isdir(DATA_CACHE):
118 logging.info(f"Loading cached resolved dataset from {DATA_CACHE}")
119 return load_from_disk(DATA_CACHE)
120
121 logging.info(f"Resolving {DATASET_NAME}/{DATASET_SUBSET} ids -> text (one-time, cached)")
122 corpus_ds = load_dataset(DATASET_NAME, "corpus", split="train")
123 corpus = dict(zip(corpus_ds["passage_id"], corpus_ds["passage"]))
124 queries_ds = load_dataset(DATASET_NAME, "queries", split="train")
125 queries = dict(zip(queries_ds["query_id"], queries_ds["query"]))
126 raw = load_dataset(DATASET_NAME, DATASET_SUBSET, split="train").select(range(TRAIN_SIZE + EVAL_SIZE))
127
128 def id_to_text(batch):
129 return {
130 "query": [queries[qid] for qid in batch["query_id"]],
131 "positive": [corpus[pid] for pid in batch["positive_id"]],
132 "negative": [corpus[pid] for pid in batch["negative_id"]],
133 "score": batch["score"],
134 }
135
136 resolved = raw.map(id_to_text, batched=True, remove_columns=["query_id", "positive_id", "negative_id"])
137 resolved.save_to_disk(DATA_CACHE)
138 return resolved
139
140
141def main() -> None:
142 parser = argparse.ArgumentParser()
143 parser.add_argument(
144 "--eval-only", type=str, default=None, help="Skip training; load this saved model and run only the evaluator."
145 )
146 cli, _ = parser.parse_known_args()
147
148 setup_logging()
149
150 if cli.eval_only:
151 logging.info(f"Eval-only mode: loading model from {cli.eval_only}")
152 model = CrossEncoder(cli.eval_only)
153 evaluator = CrossEncoderNanoBEIREvaluator(dataset_names=["msmarco", "nfcorpus", "nq"])
154 with autocast_ctx():
155 evaluator(model)
156 return
157
158 logging.info(f"Loading base model: {MODEL_NAME}")
159 model = CrossEncoder(
160 MODEL_NAME,
161 num_labels=1,
162 activation_fn=nn.Identity(), # Mandatory for distillation losses.
163 model_card_data=CrossEncoderModelCardData(
164 language="en",
165 license="apache-2.0",
166 model_name=f"{MODEL_NAME.split('/')[-1]} reranker distilled from MS MARCO ensemble",
167 ),
168 )
169
170 resolved = load_resolved_dataset()
171 if SMOKE_TEST:
172 logging.info("SMOKE_TEST=1: trimmed dataset; will run max_steps=1 and skip Hub push")
173 resolved = resolved.select(range(min(70, len(resolved))))
174 eval_size = 20 if SMOKE_TEST else EVAL_SIZE
175 split = resolved.train_test_split(test_size=eval_size, seed=12)
176 train_dataset = split["train"]
177 eval_dataset = split["test"]
178 logging.info(f" train: {len(train_dataset):,} rows | eval: {len(eval_dataset):,} rows")
179 logging.info(f" columns: {train_dataset.column_names}")
180
181 loss = MarginMSELoss(model)
182
183 evaluator = CrossEncoderNanoBEIREvaluator(dataset_names=["msmarco", "nfcorpus", "nq"])
184 logging.info("Baseline evaluation:")
185 with autocast_ctx():
186 # Must run before deriving metric_key: evaluator(model) mutates primary_metric to add the name_ prefix.
187 baseline_eval = evaluator(model)[evaluator.primary_metric]
188 metric_key = f"eval_{evaluator.primary_metric}"
189
190 args = CrossEncoderTrainingArguments(
191 output_dir=OUTPUT_DIR,
192 num_train_epochs=1,
193 max_steps=1 if SMOKE_TEST else -1,
194 per_device_train_batch_size=32,
195 per_device_eval_batch_size=32,
196 learning_rate=8e-6, # Lower than typical 2e-5. Distillation regression converges faster
197 weight_decay=0.01,
198 warmup_steps=0.1,
199 lr_scheduler_type="linear",
200 bf16=True,
201 eval_strategy="steps",
202 eval_steps=0.1,
203 save_strategy="steps",
204 save_steps=0.1,
205 save_total_limit=2,
206 logging_steps=0.01,
207 logging_first_step=True,
208 load_best_model_at_end=True,
209 metric_for_best_model=metric_key,
210 greater_is_better=True,
211 report_to="none" if SMOKE_TEST else "trackio",
212 run_name=RUN_NAME,
213 seed=12,
214 )
215
216 trainer = CrossEncoderTrainer(
217 model=model,
218 args=args,
219 train_dataset=train_dataset,
220 eval_dataset=eval_dataset,
221 loss=loss,
222 evaluator=evaluator,
223 callbacks=[EarlyStoppingCallback(early_stopping_patience=3)], # CE rerankers peak mid-training
224 )
225 if not SMOKE_TEST:
226 log_trackio_dashboard()
227 trainer.train()
228
229 logging.info("Post-training evaluation:")
230 with autocast_ctx():
231 score = evaluator(model)[evaluator.primary_metric]
232 delta = score - baseline_eval
233 verdict = "WIN" if delta >= 0.005 else "MARGINAL" if delta >= 0 else "REGRESSION"
234 logging.info(f"VERDICT: {verdict} | score={score:.4f} | baseline={baseline_eval:.4f} | delta={delta:+.4f}")
235
236 final_dir = f"{OUTPUT_DIR}/final"
237 model.save_pretrained(final_dir)
238 logging.info(f"Saved final model to {final_dir}")
239
240 if SMOKE_TEST:
241 logging.info("SMOKE_TEST=1: skipping Hub push")
242 return
243
244 try:
245 commit_url = model.push_to_hub(RUN_NAME)
246 logging.info(f"Pushed model to {commit_url.rsplit('/commit/', 1)[0]}")
247 except Exception:
248 import traceback
249
250 logging.error(f"Hub push failed:\n{traceback.format_exc()}")
251
252
253if __name__ == "__main__":
254 main()