Entrenar y afinar modelos Sentence-Transformers
Router guiado para entrenar o afinar modelos de embeddings (bi-encoder, cross-encoder, sparse encoder) con sentence-transformers, cubriendo selección de pérdidas, minería de negativos duros, evaluadores y publicación en Hugging Face Hub. Esencial para construir sistemas RAG, búsqueda semántica y reranking propios.
Incluida en el Pase · para Python, HuggingFace Hub, HuggingFace Jobs
#!/usr/bin/env python3
/// script
requires-python = ">=3.10"
dependencies = [
"sentence-transformers[train]>=5.0",
"datasets>=2.19.0",
"accelerate>=0.26.0",
"trackio",
]
///
""" LegalPilot SL — Bi-encoder Matryoshka finetuned para búsqueda de jurisprudencia española.
Cliente: LegalPilot SL (SaaS B2B para despachos de abogados) Caso: Búsqueda semántica de sentencias del CENDOJ Modelo base: intfloat/multilingual-e5-base Loss: MatryoshkaLoss(MultipleNegativesRankingLoss) → despliegue flexible 768/512/256/128 dims Datos: 42.000 pares (consulta_abogado, fragmento_jurisprudencia) — anchor/positive sin negativos explícitos Hardware: 1× NVIDIA A10G (24 GB VRAM) — Hugging Face Jobs
Ejecución: # Smoke-test (1 paso, sin push al Hub): SMOKE_TEST=1 python resultado.py
# Entrenamiento completo:
python resultado.py
# Multi-GPU:
accelerate launch resultado.py
# Hugging Face Jobs (pegar contenido del script en `script`):
hf_jobs("uv", {
"script": "<contenido de este archivo>",
"flavor": "a10g-large",
"timeout": "4h",
"secrets": {"HF_TOKEN": "$HF_TOKEN"},
})
Notas de diseño
- Prefijos E5: añadir "query: " a las consultas y "passage: " a los fragmentos es obligatorio para multilingual-e5-base. Sin ellos el modelo ignora la asimetría consulta/documento y el nDCG cae ~8 puntos. Ver references/prompts_and_instructions.md.
- MatryoshkaLoss: se entrena una sola vez pero se puede desplegar a 128/256/512/768 dims. Con dim=256 LegalPilot reduce la memoria de índice de FAISS un 67% con pérdida <2% nDCG.
- BatchSamplers.NO_DUPLICATES: crítico con MNRL — los duplicados dentro del batch generan falsos negativos que destruyen la señal de entrenamiento.
- load_best_model_at_end: el evaluador reporta nDCG@10; el trainer guarda el checkpoint con el mejor metric_key = "eval_NanoBEIR_mean_ndcg@10". """
from future import annotations
import logging import os from contextlib import nullcontext
import torch from datasets import Dataset, load_dataset
from sentence_transformers import ( SentenceTransformer, SentenceTransformerModelCardData, SentenceTransformerTrainer, SentenceTransformerTrainingArguments, ) from sentence_transformers.base.sampler import BatchSamplers from sentence_transformers.sentence_transformer.evaluation import ( InformationRetrievalEvaluator, NanoBEIREvaluator, ) from sentence_transformers.sentence_transformer.losses import ( MatryoshkaLoss, MultipleNegativesRankingLoss, )
─────────────────────────────────────────────────────────────────────────────
Configuración del run
─────────────────────────────────────────────────────────────────────────────
MODEL_NAME = "intfloat/multilingual-e5-base" DATASET_NAME = "legalpilot-sl/cendoj-pares-42k" # dataset privado en HF Hub TRAIN_SIZE = 38_000 EVAL_SIZE = 4_000 IR_EVAL_SIZE = 500 # subconjunto para InformationRetrievalEvaluator OUTPUT_DIR = "models/legalpilot-e5-matryoshka" RUN_NAME = "legalpilot-e5-matryoshka-v1" HUB_MODEL_ID = "legalpilot-sl/e5-base-jurisprudencia-es-v1" # privado en la org
MATRYOSHKA_DIMS = [768, 512, 256, 128] # dimensiones de despliegue
SMOKE_TEST = os.environ.get("SMOKE_TEST") == "1"
─────────────────────────────────────────────────────────────────────────────
Helpers
─────────────────────────────────────────────────────────────────────────────
def autocast_ctx(): """bf16/fp16 autocast para llamadas al evaluador fuera del trainer.""" if not torch.cuda.is_available(): return nullcontext() dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16 return torch.autocast("cuda", dtype=dtype)
def setup_logging() -> None: """Configura logging + TF32. Silencia spam HTTP de HuggingFace.""" os.makedirs("logs", exist_ok=True) logging.basicConfig( format="%(asctime)s - %(message)s", datefmt="%Y-%m-%d %H:%M:%S", level=logging.INFO, handlers=[ logging.StreamHandler(), logging.FileHandler(f"logs/{RUN_NAME}.log"), ], force=True, ) for noisy in ("httpx", "httpcore", "huggingface_hub", "urllib3", "filelock", "fsspec"): logging.getLogger(noisy).setLevel(logging.WARNING) if torch.cuda.is_available(): torch.set_float32_matmul_precision("high") # TF32 en Ampere+, sin pérdida de calidad
def add_e5_prefixes(example: dict) -> dict: """ multilingual-e5-base requiere prefijos asimétricos: - consultas del abogado → "query: " - fragmentos CENDOJ → "passage: " Sin esto el modelo pierde ~8 puntos nDCG@10 (ver references/prompts_and_instructions.md). """ example["anchor"] = "query: " + example["anchor"] example["positive"] = "passage: " + example["positive"] return example
def build_ir_evaluator(eval_ds: Dataset) -> InformationRetrievalEvaluator: """ Construye un InformationRetrievalEvaluator de 500 consultas contra el corpus de pasajes del split de evaluación. Métrica principal: nDCG@10. """ subset = eval_ds.select(range(min(IR_EVAL_SIZE, len(eval_ds))))
queries = {str(i): row["anchor"] for i, row in enumerate(subset)}
corpus = {str(i): row["positive"] for i, row in enumerate(subset)}
# Relevantes: 1 documento correcto por consulta (el emparejado)
relevant = {str(i): {str(i)} for i in range(len(subset))}
return InformationRetrievalEvaluator(
queries=queries,
corpus=corpus,
relevant_docs=relevant,
name="legalpilot-eval",
score_functions={"cosine": "cos_sim"},
show_progress_bar=False,
)
─────────────────────────────────────────────────────────────────────────────
Main
─────────────────────────────────────────────────────────────────────────────
def main() -> None: setup_logging() logging.info(f"Run: {RUN_NAME}") logging.info(f"Modelo base: {MODEL_NAME}") logging.info(f"Dims Matryoshka: {MATRYOSHKA_DIMS}")
# ── 1. Cargar modelo ──────────────────────────────────────────────────────
model = SentenceTransformer(
MODEL_NAME,
model_card_data=SentenceTransformerModelCardData(
language="es",
license="apache-2.0",
model_name="LegalPilot E5 Matryoshka — Jurisprudencia Española",
),
)
# ── 2. Cargar y preparar datos ────────────────────────────────────────────
train_size = 50 if SMOKE_TEST else TRAIN_SIZE
eval_size = 20 if SMOKE_TEST else EVAL_SIZE
logging.info(f"Cargando dataset: {DATASET_NAME}")
raw = load_dataset(DATASET_NAME, split="train") # el dataset tiene un único split
raw = raw.map(add_e5_prefixes) # añadir prefijos E5
train_ds = raw.select(range(train_size))
eval_ds = raw.select(range(train_size, train_size + eval_size))
logging.info(f" train: {len(train_ds):,} ejemplos")
logging.info(f" eval: {len(eval_ds):,} ejemplos")
# ── 3. Loss: MatryoshkaLoss(MNRL) ─────────────────────────────────────────
base_loss = MultipleNegativesRankingLoss(model)
loss = MatryoshkaLoss(
model,
loss=base_loss,
matryoshka_dims=MATRYOSHKA_DIMS,
)
# ── 4. Evaluador ──────────────────────────────────────────────────────────
evaluator = build_ir_evaluator(eval_ds)
metric_key = f"eval_{evaluator.primary_metric}"
logging.info(f"metric_for_best_model: {metric_key}")
# ── 5. Evaluación baseline (antes de entrenar) ────────────────────────────
logging.info("Evaluación baseline (antes del entrenamiento):")
with autocast_ctx():
baseline_eval = evaluator(model)[evaluator.primary_metric]
logging.info(f" baseline {metric_key} = {baseline_eval:.4f}")
# ── 6. Training arguments ─────────────────────────────────────────────────
args = SentenceTransformerTrainingArguments(
output_dir=OUTPUT_DIR,
num_train_epochs=3,
max_steps=1 if SMOKE_TEST else -1,
per_device_train_batch_size=128,
per_device_eval_batch_size=128,
learning_rate=2e-5,
weight_decay=0.01,
warmup_steps=0.1, # 10% del total de pasos (float, no deprecated warmup_ratio)
lr_scheduler_type="linear",
bf16=True, # carga fp32 + autocast bf16; nunca torch_dtype=bfloat16
batch_sampler=BatchSamplers.NO_DUPLICATES, # crítico para MNRL
eval_strategy="steps",
eval_steps=0.1,
save_strategy="steps",
save_steps=0.1, # múltiplo de eval_steps → load_best_model_at_end funciona
save_total_limit=2,
logging_steps=0.01,
logging_first_step=True,
load_best_model_at_end=True,
metric_for_best_model=metric_key,
greater_is_better=True,
# HF Jobs: activar push in-trainer para entorno efímero
push_to_hub=True,
hub_model_id=HUB_MODEL_ID,
hub_strategy="every_save",
hub_private_repo=True,
report_to="none" if SMOKE_TEST else "trackio",
run_name=RUN_NAME,
seed=42,
)
# ── 7. Entrenamiento ──────────────────────────────────────────────────────
trainer = SentenceTransformerTrainer(
model=model,
args=args,
train_dataset=train_ds,
eval_dataset=eval_ds,
loss=loss,
evaluator=evaluator,
)
if not SMOKE_TEST:
try:
from huggingface_hub import whoami
hf_user = whoami().get("name")
if hf_user:
logging.info(
f"Dashboard Trackio (progreso en vivo): "
f"https://huggingface.co/spaces/{hf_user}/trackio"
)
except Exception:
pass
trainer.train()
# ── 8. Evaluación post-entrenamiento y veredicto ──────────────────────────
logging.info("Evaluación post-entrenamiento:")
with autocast_ctx():
score = evaluator(model)[evaluator.primary_metric]
delta = score - baseline_eval
verdict = "WIN" if delta >= 0.005 else "MARGINAL" if delta >= 0 else "REGRESSION"
logging.info(
f"VERDICT: {verdict} | score={score:.4f} | baseline={baseline_eval:.4f} | delta={delta:+.4f}"
)
# ── 9. Guardar modelo final ───────────────────────────────────────────────
final_dir = f"{OUTPUT_DIR}/final"
model.save_pretrained(final_dir)
logging.info(f"Modelo guardado en: {final_dir}")
# ── 10. Push al Hub (privado, org legalpilot-sl) ──────────────────────────
if SMOKE_TEST:
logging.info("SMOKE_TEST=1: omitiendo push final al Hub")
return
try:
commit_url = model.push_to_hub(
HUB_MODEL_ID,
private=True,
exist_ok=True,
)
logging.info(f"Modelo publicado en: {commit_url.rsplit('/commit/', 1)[0]}")
except Exception:
import traceback
logging.error(f"Hub push falló:\n{traceback.format_exc()}")
# ── 11. Registrar experimento ─────────────────────────────────────────────
os.makedirs("logs", exist_ok=True)
with open("logs/experiments.md", "a") as f:
f.write(
f"\n## {RUN_NAME}\n"
f"- Fecha: 2026-06-18\n"
f"- Base: `{MODEL_NAME}`\n"
f"- Loss: `MatryoshkaLoss(MNRL)` dims={MATRYOSHKA_DIMS}\n"
f"- Dataset: `{DATASET_NAME}` ({TRAIN_SIZE:,} train / {EVAL_SIZE:,} eval)\n"
f"- Epochs: 3 | batch: 128 | LR: 2e-5 | warmup: 10%\n"
f"- Baseline: `{baseline_eval:.4f}` → Final: `{score:.4f}` "
f"(Δ={delta:+.4f})\n"
f"- **VERDICT: {verdict}**\n"
f"- Hub: `{HUB_MODEL_ID}`\n"
)
logging.info("Experimento registrado en logs/experiments.md")
if name == "main": main()
// qué_hace
Guia al agente para entrenar modelos de embeddings y rerankers personalizados con sentence-transformers, cubriendo los tres tipos de encoder y sus variantes avanzadas.
// cómo_lo_hace
Actua como router que identifica el tipo de modelo, carga las referencias y plantillas de produccion correctas, y aplica un workflow estructurado con smoke-test, evaluacion baseline y publicacion automatica en el Hub.
// ejemplo_de_uso
Úsala cuando necesites crear un modelo de búsqueda semántica adaptado a tu dominio, sin depender de embeddings genéricos. Ej.: entrenas un bi-encoder con tus FAQs y lo publicas en el Hub con evaluación automática incluida.
// plataformas
// opiniones_de_la_comunidad
Opiniones
Cargando opiniones…