← Volver al catálogo
IA, ingeniería y MLOpsSistemaAvanzadoEn el pase

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.

Comprobando acceso…

Incluida en el Pase · para Python, HuggingFace Hub, HuggingFace Jobs

// resultado_de_ejemplo

#!/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

PythonHuggingFace HubHuggingFace JobsGPU/CPU local
Categoría
IA, ingeniería y MLOps
Tipo
Sistema
Nivel
Avanzado
Licencia
Apache-2.0
Seguridad
seguro · riesgo alto
Versión
1.0.0

// opiniones_de_la_comunidad

Opiniones

Cargando opiniones…