Setting the file. One moment. Train Sentence Transformer With Lora Example · Train Sentence Transformers · huggingface/skills · Skills DocsLosses Cross Encoder
scripts/train_sentence_transformer_with_lora_example.py
Python·275 lines·10 KB
of the box with sentence-transformers via `peft`.
17
18Use LoRA when: the base is a large decoder (Qwen3, Llama, Mistral, Gemma) and
19a full fine-tune is VRAM-prohibitive. You want multiple task-specific adapters
20on one base. You're nudging an existing strong retriever (E5-mistral,
21Qwen3-Embedding) on domain data. Skip LoRA for small encoders (BERT-base,
22MiniLM: full fine-tune is tractable and usually better) or tiny datasets
23(<1k pairs: adapter rank becomes the bottleneck).
24
25This template defaults to BERT-base for portability. Swap `MODEL_NAME` to a
26decoder backbone for the use case LoRA actually shines on.
27
28Architecture variants:
29- Bi-encoder (this script): `TaskType.FEATURE_EXTRACTION`.
30- Sparse-encoder: same pattern. `SparseEncoder` supports `add_adapter`. See
31 `examples/sparse_encoder/training/peft/train_splade_gooaq_peft.py`.
32- Cross-encoder: use `TaskType.SEQ_CLS` for `num_labels >= 1`. Community
33 examples are sparse. Smoke-test with `max_steps=1` first.
34
35Key hyperparameters:
36- `r` (rank): 8-128. Bigger = more capacity + memory. 64 is a strong default.
37- `lora_alpha`: typically 2 x r (some teams use 1 x r for stability).
38- `lora_dropout`: 0.05-0.1. Raise to 0.1 for small datasets.
39- `target_modules=None` auto-picks attention modules. Pass
40 `["q_proj", "k_proj", "v_proj", "o_proj"]` (attention) or
41 `["gate_proj", "up_proj", "down_proj"]` (MLP) for explicit control.
42- `modules_to_save=["pooler"]` for CLS-pooled bases. The pooler Dense should
43 be trained too, not adapted.
44- LR is HIGHER than full fine-tune: 1e-4 to 5e-4 for LoRA vs. 2e-5 full.
45
46Rough memory savings on a 0.6B base (bf16, batch 64, seq 128): full fine-tune
47~24 GB, LoRA r=64 ~10 GB (~12M trainable, 2%), LoRA r=16 ~8 GB (~3M, 0.5%).
48Bigger savings on 7B+ models.
49
50Saving / sharing: `model.save_pretrained("dir")` writes ONLY the adapter (few
51MB) plus a reference to the base model. Loaders call the same one-liner.
52`peft` is invoked and the base downloaded on demand. For a merged model that
53loads without `peft` (needed for vLLM-style servers), call
54`model.transformers_model.merge_and_unload()` then `save_pretrained` /
55`push_to_hub`.
56
57Swapping adapters at inference (the main multi-task deployment win):
58 model = SentenceTransformer("base-model")
59 model.load_adapter("adapter-a", adapter_name="a")
60 model.load_adapter("adapter-b", adapter_name="b")
61 model.set_adapter("a"); emb_a = model.encode([...])
62
63QLoRA (4-bit base + LoRA) for 7B+ on consumer GPUs:
64 from transformers import BitsAndBytesConfig
65 bnb = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_quant_type="nf4",
66 bnb_4bit_compute_dtype=torch.bfloat16)
67 model = SentenceTransformer("Qwen/Qwen3-Embedding-7B",
68 model_kwargs={"quantization_config": bnb})
69 model.add_adapter(LoraConfig(r=64, lora_alpha=128, ...))
70`pip install bitsandbytes` first. Linux-only for bitsandbytes (Windows: the
71fork or WSL).
72
73Known issues:
74- LoRA PEFT on Qwen2.5-VL / paligemma / gemma3 / internvl / aya_vision under
75 transformers v5: `AutoModel.from_pretrained(peft_path)` crashes with
76 `KeyError: 'qwen2_vl'`. Pin transformers to 4.x or wait for the upstream fix.
77- `gradient_checkpointing=True` + LoRA: usually works. If you hit "None of
78 the inputs have requires_grad=True", call
79 `model.transformers_model.enable_input_require_grads()` after `add_adapter`.
80- `add_adapter` before pooling: when building from scratch (not loading a
81 pre-assembled checkpoint), call `add_adapter` AFTER
82 `SentenceTransformer(modules=[...])` is complete.
83
84Common gotchas: LR still at 2e-5 (LoRA needs higher). Forgetting to merge for
85vLLM-style servers (they don't load `peft`). `r=8` too small for retrievers
86trained on millions of pairs (try 32 or 64). `modules_to_save` missing the
87pooler on CLS-pooled bases.
88"""
89
90from __future__ import annotations
91
92import argparse
93import logging
94import os
95from contextlib import nullcontext
96
97import torch
98from datasets import load_dataset
99from peft import LoraConfig, TaskType
100
101from sentence_transformers import (
102 SentenceTransformer,
103 SentenceTransformerModelCardData,
104 SentenceTransformerTrainer,
105 SentenceTransformerTrainingArguments,
106)
107from sentence_transformers.base.sampler import BatchSamplers
108from sentence_transformers.sentence_transformer.evaluation import NanoBEIREvaluator
109from sentence_transformers.sentence_transformer.losses import CachedMultipleNegativesRankingLoss
110
111
112def autocast_ctx():
113 """bf16/fp16 autocast for evaluator calls outside the trainer (which has its own autocast)."""
114 if not torch.cuda.is_available():
115 return nullcontext()
116 dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
117 return torch.autocast("cuda", dtype=dtype)
118
119
120def log_trackio_dashboard():
121 """Surface the Trackio dashboard URL so the user can watch training live."""
122 try:
123 from huggingface_hub import whoami
124
125 hf_user = whoami().get("name")
126 if hf_user:
127 logging.info(
128 f"Trackio dashboard (live training progress): https://huggingface.co/spaces/{hf_user}/trackio"
129 )
130 except Exception:
131 pass
132
133
134MODEL_NAME = "google-bert/bert-base-uncased"
135OUTPUT_DIR = "models/bert-base-gooaq-lora"
136RUN_NAME = "bert-base-gooaq-lora"
137SMOKE_TEST = os.environ.get("SMOKE_TEST") == "1"
138
139
140def setup_logging():
141 """Configure logging + TF32. Tees to logs/{RUN_NAME}.log and silences HTTP spam."""
142 os.makedirs("logs", exist_ok=True)
143 logging.basicConfig(
144 format="%(asctime)s - %(message)s",
145 datefmt="%Y-%m-%d %H:%M:%S",
146 level=logging.INFO,
147 handlers=[logging.StreamHandler(), logging.FileHandler(f"logs/{RUN_NAME}.log")],
148 force=True,
149 )
150 for noisy in ("httpx", "httpcore", "huggingface_hub", "urllib3", "filelock", "fsspec"):
151 logging.getLogger(noisy).setLevel(logging.WARNING)
152 if torch.cuda.is_available():
153 torch.set_float32_matmul_precision("high") # TF32 on Ampere+, no quality loss
154
155
156def main() -> None:
157 parser = argparse.ArgumentParser()
158 parser.add_argument(
159 "--eval-only", type=str, default=None, help="Skip training; load this saved model and run only the evaluator."
160 )
161 cli, _ = parser.parse_known_args()
162
163 setup_logging()
164
165 if cli.eval_only:
166 logging.info(f"Eval-only mode: loading model from {cli.eval_only}")
167 model = SentenceTransformer(cli.eval_only)
168 evaluator = NanoBEIREvaluator()
169 with autocast_ctx():
170 evaluator(model)
171 return
172
173 model = SentenceTransformer(
174 MODEL_NAME,
175 model_card_data=SentenceTransformerModelCardData(
176 language="en",
177 license="apache-2.0",
178 model_name=f"{MODEL_NAME.split('/')[-1]} LoRA adapter on GooAQ",
179 ),
180 )
181
182 peft_config = LoraConfig(
183 task_type=TaskType.FEATURE_EXTRACTION,
184 inference_mode=False,
185 r=64,
186 lora_alpha=128,
187 lora_dropout=0.1,
188 )
189 model.add_adapter(peft_config)
190 trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
191 total = sum(p.numel() for p in model.parameters())
192 logging.info(f"trainable params: {trainable:,} / {total:,} ({100 * trainable / total:.2f}%)")
193
194 full = load_dataset("sentence-transformers/gooaq", split="train")
195 if SMOKE_TEST:
196 logging.info("SMOKE_TEST=1: trimmed dataset; will run max_steps=1 and skip Hub push")
197 full = full.select(range(min(200, len(full))))
198 eval_size = 20 if SMOKE_TEST else 10_000
199 train_cap = 50 if SMOKE_TEST else 500_000
200 split = full.train_test_split(test_size=eval_size, seed=12)
201 train_dataset = split["train"].select(range(min(train_cap, len(split["train"]))))
202 eval_dataset = split["test"]
203 logging.info(f"train={len(train_dataset):,} eval={len(eval_dataset):,}")
204
205 loss = CachedMultipleNegativesRankingLoss(model, mini_batch_size=32)
206
207 evaluator = NanoBEIREvaluator()
208 logging.info("Baseline:")
209 with autocast_ctx():
210 # Must run before deriving metric_key: evaluator(model) mutates primary_metric to add the name_ prefix.
211 baseline_eval = evaluator(model)[evaluator.primary_metric]
212 metric_key = f"eval_{evaluator.primary_metric}"
213
214 args = SentenceTransformerTrainingArguments(
215 output_dir=OUTPUT_DIR,
216 num_train_epochs=1,
217 max_steps=1 if SMOKE_TEST else -1,
218 per_device_train_batch_size=512,
219 per_device_eval_batch_size=512,
220 learning_rate=1e-4,
221 weight_decay=0.01,
222 warmup_steps=0.1,
223 bf16=True,
224 batch_sampler=BatchSamplers.NO_DUPLICATES,
225 eval_strategy="steps",
226 eval_steps=0.1,
227 save_strategy="steps",
228 save_steps=0.1,
229 save_total_limit=2,
230 logging_steps=0.01,
231 logging_first_step=True,
232 load_best_model_at_end=True,
233 metric_for_best_model=metric_key,
234 greater_is_better=True,
235 report_to="none" if SMOKE_TEST else "trackio",
236 run_name=RUN_NAME,
237 seed=12,
238 )
239
240 trainer = SentenceTransformerTrainer(
241 model=model,
242 args=args,
243 train_dataset=train_dataset,
244 eval_dataset=eval_dataset,
245 loss=loss,
246 evaluator=evaluator,
247 )
248 if not SMOKE_TEST:
249 log_trackio_dashboard()
250 trainer.train()
251
252 logging.info("Post-training evaluation:")
253 with autocast_ctx():
254 score = evaluator(model)[evaluator.primary_metric]
255 delta = score - baseline_eval
256 verdict = "WIN" if delta >= 0.005 else "MARGINAL" if delta >= 0 else "REGRESSION"
257 logging.info(f"VERDICT: {verdict} | score={score:.4f} | baseline={baseline_eval:.4f} | delta={delta:+.4f}")
258
259 model.save_pretrained(f"{OUTPUT_DIR}/final")
260
261 if SMOKE_TEST:
262 logging.info("SMOKE_TEST=1: skipping Hub push")
263 return
264
265 try:
266 commit_url = model.push_to_hub(RUN_NAME)
267 logging.info(f"Pushed model to {commit_url.rsplit('/commit/', 1)[0]}")
268 except Exception:
269 import traceback
270
271 logging.error(f"Hub push failed:\n{traceback.format_exc()}")
272
273
274if __name__ == "__main__":
275 main()