Generación de datos sintéticos con la Batch API

Etapa 3 del plan. Se elige la Batch API sobre llamadas sueltas porque cuesta
la mitad y generar un dataset no es sensible a la latencia: es exactamente su
caso de uso. A cambio el trabajo es asincrónico y puede tardar horas, así que
todo el estado vive en disco y es reanudable — matar el proceso, perder la
conexión o reiniciar el servidor no pierde nada ni gasta de nuevo.

Cuatro decisiones que sostienen eso:

- El custom_id se deriva del hash del contenido del pedido. El trabajo queda
  idempotente: re-generar la misma especificación produce los mismos IDs, así
  que lo ya resuelto se reconoce y no se vuelve a pagar.
- Los resultados se indexan por custom_id, nunca por posición. La API los
  devuelve en cualquier orden, y emparejarlos por índice mezcla las respuestas
  en silencio; un dataset mal alineado es peor que uno vacío porque parece
  correcto. El cliente falso de los tests los invierte a propósito.
- Los fallos se registran en errors.jsonl con su motivo en vez de descartarse,
  y --retry-failed reenvía solo esos.
- El prompt de sistema compartido va marcado para caché: es idéntico entre
  miles de pedidos, y a ~0,1x del precio de entrada deja de contar.

Los lotes se persisten después de cada envío y no al final, para que una caída
a mitad de camino no deje lotes huérfanos sin registrar.

El SDK de anthropic se importa de forma perezosa y queda como extra opcional:
el servidor de entrenamiento no lo necesita y el paquete importa sin él.

21 tests nuevos (92 en total), todos contra un cliente falso — lo que hay que
verificar es el harness, no el SDK.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-27 23:26:41 -03:00
parent 1d2fcef4d0
commit 8805b0541b
5 changed files with 419 additions and 18 deletions
+9 -1
View File
@@ -4,11 +4,17 @@ Todo hiperparámetro, ruta y umbral vive en `configs/*.yaml`, se compone por
capas y se valida con pydantic antes de que arranque cualquier proceso largo. capas y se valida con pydantic antes de que arranque cualquier proceso largo.
""" """
from enlace.config.load import load_agent_config, load_config, load_config_from_argv from enlace.config.load import (
load_agent_config,
load_config,
load_config_from_argv,
load_distill_config,
)
from enlace.config.schema import ( from enlace.config.schema import (
AgentConfig, AgentConfig,
Config, Config,
DataConfig, DataConfig,
DistillConfig,
HardwareConfig, HardwareConfig,
ModelConfig, ModelConfig,
OptimizerConfig, OptimizerConfig,
@@ -21,6 +27,7 @@ __all__ = [
"AgentConfig", "AgentConfig",
"Config", "Config",
"DataConfig", "DataConfig",
"DistillConfig",
"HardwareConfig", "HardwareConfig",
"ModelConfig", "ModelConfig",
"OptimizerConfig", "OptimizerConfig",
@@ -30,4 +37,5 @@ __all__ = [
"load_agent_config", "load_agent_config",
"load_config", "load_config",
"load_config_from_argv", "load_config_from_argv",
"load_distill_config",
] ]
+38 -17
View File
@@ -15,7 +15,7 @@ from typing import Any
from omegaconf import DictConfig, OmegaConf from omegaconf import DictConfig, OmegaConf
from enlace.config.schema import AgentConfig, Config from enlace.config.schema import AgentConfig, Config, DistillConfig
# Familias de capas que puede declarar un YAML raíz, en el orden en que se # Familias de capas que puede declarar un YAML raíz, en el orden en que se
# fusionan. El orden solo importa para los mensajes de error. # fusionan. El orden solo importa para los mensajes de error.
@@ -100,19 +100,17 @@ def load_config(path: str | Path, overrides: list[str] | None = None) -> Config:
raise ConfigError(f"config inválida ({path}):\n{exc}") from None raise ConfigError(f"config inválida ({path}):\n{exc}") from None
def load_agent_config( CONFIGS_ROOT = Path(__file__).resolve().parents[2] / "configs"
path: str | Path | None = None, overrides: list[str] | None = None
) -> AgentConfig:
"""Carga la config del runtime del agente.
Es un árbol aparte del de entrenamiento: el agente no necesita saber nada
del optimizador ni del schedule. Los valores que salen de `.env` se resuelven def _load_single(path: Path, overrides: list[str] | None, model_cls, etiqueta: str):
con `${oc.env:...}`, así que los secretos y las URLs de la instalación nunca """Carga un YAML suelto y lo valida. No compone capas.
entran en un archivo commiteado.
Lo usan los árboles de config que no son de entrenamiento (agente,
destilación): no necesitan saber nada del optimizador ni del schedule.
Los valores que salen de `.env` se resuelven con `${oc.env:...}`, así que
los secretos nunca entran en un archivo commiteado.
""" """
if path is None:
path = Path(__file__).resolve().parents[2] / "configs" / "agent" / "tools.yaml"
path = Path(path)
if not path.is_file(): if not path.is_file():
raise ConfigError(f"no existe el archivo de config: {path}") raise ConfigError(f"no existe el archivo de config: {path}")
@@ -124,14 +122,37 @@ def load_agent_config(
# OmegaConf devuelve cadena vacía cuando la variable de entorno no está # OmegaConf devuelve cadena vacía cuando la variable de entorno no está
# definida; para pydantic eso tiene que ser ausencia, no un valor vacío. # definida; para pydantic eso tiene que ser ausencia, no un valor vacío.
search = data.get("search") _vaciar_a_none(data)
if isinstance(search, dict) and search.get("searxng_url") == "":
search["searxng_url"] = None
try: try:
return AgentConfig.model_validate(data) return model_cls.model_validate(data)
except Exception as exc: except Exception as exc:
raise ConfigError(f"config de agente inválida ({path}):\n{exc}") from None raise ConfigError(f"config de {etiqueta} inválida ({path}):\n{exc}") from None
def _vaciar_a_none(data: dict) -> None:
"""Convierte las cadenas vacías de `${oc.env:VAR,""}` en ausencia, recursivo."""
for clave, valor in data.items():
if isinstance(valor, dict):
_vaciar_a_none(valor)
elif valor == "":
data[clave] = None
def load_agent_config(
path: str | Path | None = None, overrides: list[str] | None = None
) -> AgentConfig:
"""Carga la config del runtime del agente."""
path = Path(path) if path is not None else CONFIGS_ROOT / "agent" / "tools.yaml"
return _load_single(path, overrides, AgentConfig, "agente")
def load_distill_config(
path: str | Path | None = None, overrides: list[str] | None = None
) -> DistillConfig:
"""Carga la config de generación de datos sintéticos por lotes."""
path = Path(path) if path is not None else CONFIGS_ROOT / "data" / "distill.yaml"
return _load_single(path, overrides, DistillConfig, "destilación")
def load_config_from_argv(argv: list[str] | None = None) -> Config: def load_config_from_argv(argv: list[str] | None = None) -> Config:
+34
View File
@@ -242,6 +242,40 @@ class AgentConfig(_Base):
search: SearchConfig = SearchConfig() search: SearchConfig = SearchConfig()
class DistillConfig(_Base):
"""Generación de datos sintéticos por lotes (Etapa 3 del plan).
Se usa la Batch API y no llamadas sueltas por dos razones: cuesta la mitad,
y generar un dataset no es sensible a la latencia — es exactamente su caso
de uso. A cambio hay que tratar el trabajo como asincrónico, con estado en
disco y reanudable, que es lo que hace `enlace/data/distill.py`.
"""
model: str = "claude-opus-5"
# En este modelo el pensamiento está activo por defecto y `max_tokens` acota
# pensamiento + respuesta juntos: sin holgura la salida se corta a la mitad.
max_tokens: int = Field(default=8000, gt=0)
effort: Literal["low", "medium", "high", "xhigh", "max"] = "low"
# Límites de la Batch API: 100.000 pedidos o 256 MB por lote. Se parte en
# trozos bien por debajo del tope para que un trabajo grande no dependa de
# que una sola request HTTP enorme llegue entera.
requests_per_batch: int = Field(default=10_000, gt=0, le=100_000)
poll_interval_seconds: float = Field(default=60.0, ge=5.0)
# La API promete resultados en menos de 24 h; el tope propio es una red de
# seguridad para que un proceso olvidado no quede colgado para siempre.
max_wait_hours: float = Field(default=26.0, gt=0)
output_dir: str = "data/distill"
# Solo para el reporte de costo estimado; los precios reales los factura
# Anthropic. Se declaran acá para no hardcodearlos en el código.
usd_per_mtok_input: float = Field(default=5.0, ge=0)
usd_per_mtok_output: float = Field(default=25.0, ge=0)
batch_discount: float = Field(default=0.5, gt=0, le=1.0)
class Config(_Base): class Config(_Base):
"""Config raíz: la composición de todas las capas.""" """Config raíz: la composición de todas las capas."""
+6
View File
@@ -21,6 +21,12 @@ data = [
"fasttext-wheel>=0.9", # filtro de idioma "fasttext-wheel>=0.9", # filtro de idioma
"datasketch>=1.6", # deduplicación MinHash "datasketch>=1.6", # deduplicación MinHash
] ]
distill = [
# Solo para generar el dataset SFT (Etapa 3) con la Batch API. No hace falta
# en el servidor de entrenamiento: enlace/data/distill.py importa el SDK de
# forma perezosa para que el paquete funcione sin él.
"anthropic>=0.69",
]
dev = [ dev = [
"pytest>=8.0", "pytest>=8.0",
"ruff>=0.6", "ruff>=0.6",
+332
View File
@@ -0,0 +1,332 @@
"""Generación por lotes: idempotencia, reanudación y emparejado por custom_id.
Todo corre contra un cliente falso. Lo que hay que verificar es el harness —
que no duplique costo, que no pierda resultados y que no desalinee el dataset —
no el SDK de Anthropic.
"""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any, Iterator
import pytest
from enlace.config.load import load_distill_config
from enlace.config.schema import DistillConfig
from enlace.data.distill import (
DistillError,
Peticion,
Trabajo,
collect,
leer_especificaciones,
resumen_de_uso,
status,
submit,
)
class ClienteFalso:
"""Cliente de lotes en memoria, con fallos y demoras programables."""
def __init__(
self,
*,
fallan: set[str] | None = None,
rondas_hasta_terminar: int = 0,
desordenar: bool = True,
) -> None:
self.lotes: dict[str, list[dict[str, Any]]] = {}
self.fallan = fallan or set()
self.rondas_hasta_terminar = rondas_hasta_terminar
self.desordenar = desordenar
self.consultas = 0
self.envios = 0
def construir_params(self, peticion: Peticion) -> dict[str, Any]:
return {"model": "falso", "user": peticion.user}
def crear(self, pedidos: list[dict[str, Any]]) -> str:
self.envios += 1
batch_id = f"batch-{self.envios:03d}"
self.lotes[batch_id] = pedidos
return batch_id
def estado(self, batch_id: str) -> dict[str, Any]:
self.consultas += 1
listo = self.consultas > self.rondas_hasta_terminar * len(self.lotes)
return {
"processing_status": "ended" if listo else "in_progress",
"succeeded": len(self.lotes[batch_id]),
"errored": 0,
"processing": 0,
"canceled": 0,
"expired": 0,
}
def resultados(self, batch_id: str) -> Iterator[dict[str, Any]]:
pedidos = self.lotes[batch_id]
# La API devuelve los resultados en cualquier orden: el doble los
# invierte a propósito para que un emparejado por posición falle.
if self.desordenar:
pedidos = list(reversed(pedidos))
for pedido in pedidos:
cid = pedido["custom_id"]
if cid in self.fallan:
yield {"custom_id": cid, "type": "errored", "error_type": "api_error"}
else:
yield {
"custom_id": cid,
"type": "succeeded",
"text": f"respuesta para {pedido['params']['user']}",
"stop_reason": "end_turn",
"usage": {
"input_tokens": 100,
"output_tokens": 200,
"cache_read_input_tokens": 900,
"cache_creation_input_tokens": 0,
},
}
def cancelar(self, batch_id: str) -> None:
self.lotes.pop(batch_id, None)
@pytest.fixture
def config() -> DistillConfig:
return DistillConfig(requests_per_batch=3, poll_interval_seconds=5.0)
@pytest.fixture
def trabajo(tmp_path: Path) -> Trabajo:
return Trabajo(nombre="prueba", raiz=tmp_path)
def _peticiones(n: int) -> list[Peticion]:
return [
Peticion(system="Sos ENLACE.", user=f"consulta {i}", meta={"semilla": i})
for i in range(n)
]
def _sin_espera(trabajo, cliente, config, **kw):
return collect(trabajo, cliente, config, ahora=lambda: 0.0, dormir=lambda _: None, **kw)
# --- custom_id -------------------------------------------------------------
def test_el_custom_id_se_deriva_del_contenido():
"""Idempotencia: la misma especificación produce el mismo identificador,
así que re-generar un trabajo no vuelve a pagar lo ya resuelto."""
a = Peticion(system="S", user="U")
b = Peticion(system="S", user="U", meta={"otra": "cosa"})
assert a.custom_id == b.custom_id # los metadatos no van a la API
c = Peticion(system="S", user="OTRA")
assert a.custom_id != c.custom_id
def test_las_peticiones_duplicadas_se_descartan(trabajo):
guardadas = trabajo.guardar_peticiones([*_peticiones(3), *_peticiones(3)])
assert len(guardadas) == 3
# --- submit ----------------------------------------------------------------
def test_submit_parte_en_lotes_del_tamano_configurado(trabajo, config):
cliente = ClienteFalso()
ids = submit(trabajo, _peticiones(7), cliente, config)
assert len(ids) == 3 # 3 + 3 + 1
assert cliente.envios == 3
def test_submit_persiste_cada_lote_al_enviarlo(trabajo, config):
"""Si el proceso muere en el segundo envío, el primero no queda huérfano."""
class ClienteQueFalla(ClienteFalso):
def crear(self, pedidos):
if self.envios >= 1:
raise RuntimeError("caída simulada")
return super().crear(pedidos)
with pytest.raises(RuntimeError):
submit(trabajo, _peticiones(7), ClienteQueFalla(), config)
assert len(trabajo.leer_lotes()) == 1 # el primero quedó registrado
def test_submit_no_reenvia_lo_ya_enviado(trabajo, config):
cliente = ClienteFalso()
submit(trabajo, _peticiones(6), cliente, config)
submit(trabajo, _peticiones(6), cliente, config)
assert cliente.envios == 2 # los mismos dos lotes, no cuatro
def test_submit_no_reenvia_lo_ya_resuelto(trabajo, config):
cliente = ClienteFalso()
submit(trabajo, _peticiones(6), cliente, config)
_sin_espera(trabajo, cliente, config)
# Un trabajo ampliado: 6 viejas ya resueltas + 2 nuevas.
cliente2 = ClienteFalso()
submit(trabajo, _peticiones(8), cliente2, config)
enviados = sum(len(p) for p in cliente2.lotes.values())
assert enviados == 2
# --- collect ---------------------------------------------------------------
def test_collect_empareja_por_custom_id_no_por_posicion(trabajo, config):
"""El cliente falso devuelve los resultados invertidos a propósito.
Emparejar por índice mezclaría las respuestas en silencio, y un dataset mal
alineado es peor que uno vacío porque parece correcto.
"""
cliente = ClienteFalso(desordenar=True)
peticiones = _peticiones(3)
submit(trabajo, peticiones, cliente, config)
_sin_espera(trabajo, cliente, config)
por_id = {}
with trabajo.resultados_path.open(encoding="utf-8") as fh:
for linea in fh:
fila = json.loads(linea)
por_id[fila["custom_id"]] = fila["text"]
for p in peticiones:
assert por_id[p.custom_id] == f"respuesta para {p.user}"
def test_los_fallos_van_a_errors_y_no_se_pierden(trabajo, config):
peticiones = _peticiones(4)
cliente = ClienteFalso(fallan={peticiones[1].custom_id})
submit(trabajo, peticiones, cliente, config)
resumen = _sin_espera(trabajo, cliente, config)
assert resumen["ok"] == 3
assert resumen["error"] == 1
errores = trabajo.errores_path.read_text().strip().splitlines()
assert json.loads(errores[0])["custom_id"] == peticiones[1].custom_id
def test_recoger_dos_veces_no_duplica(trabajo, config):
cliente = ClienteFalso()
submit(trabajo, _peticiones(4), cliente, config)
primero = _sin_espera(trabajo, cliente, config)
segundo = _sin_espera(trabajo, cliente, config)
assert primero["ok"] == 4
assert segundo["ok"] == 0
assert segundo["ya_estaban"] == 4
assert len(trabajo.resultados_path.read_text().strip().splitlines()) == 4
def test_collect_espera_a_los_lotes_en_curso(trabajo, config):
cliente = ClienteFalso(rondas_hasta_terminar=2)
submit(trabajo, _peticiones(3), cliente, config)
dormidas = []
resumen = collect(
trabajo, cliente, config, ahora=lambda: 0.0, dormir=dormidas.append
)
assert resumen["ok"] == 3
assert dormidas # efectivamente esperó
def test_collect_respeta_el_tope_de_espera(trabajo, config):
cliente = ClienteFalso(rondas_hasta_terminar=1000)
submit(trabajo, _peticiones(3), cliente, config)
reloj = iter([0.0, 1e9, 1e9])
with pytest.raises(DistillError, match="max_wait_hours"):
collect(
trabajo, cliente, config, ahora=lambda: next(reloj), dormir=lambda _: None
)
def test_collect_sin_lotes_da_un_error_util(trabajo, config):
with pytest.raises(DistillError, match="submit"):
_sin_espera(trabajo, ClienteFalso(), config)
def test_status_reporta_cada_lote(trabajo, config):
cliente = ClienteFalso()
submit(trabajo, _peticiones(7), cliente, config)
estados = status(trabajo, cliente)
assert len(estados) == 3
assert all("processing_status" in e for e in estados)
# --- costo -----------------------------------------------------------------
def test_el_resumen_de_uso_aplica_el_descuento_y_la_cache(trabajo, config):
cliente = ClienteFalso()
submit(trabajo, _peticiones(2), cliente, config)
_sin_espera(trabajo, cliente, config)
uso = resumen_de_uso(trabajo, config)
assert uso["resultados"] == 2
assert uso["output_tokens"] == 400
assert uso["cache_read_input_tokens"] == 1800
# (200 entrada + 1800*0.1 caché)/1e6*5 + 400/1e6*25, todo al 50%.
esperado = ((200 + 180) / 1e6 * 5 + 400 / 1e6 * 25) * 0.5
assert uso["usd_estimado"] == pytest.approx(round(esperado, 2))
def test_sin_resultados_el_costo_es_cero(trabajo, config):
assert resumen_de_uso(trabajo, config)["usd_estimado"] == 0.0
# --- entrada ---------------------------------------------------------------
def test_lee_especificaciones_jsonl(tmp_path):
path = tmp_path / "peticiones.jsonl"
path.write_text(
json.dumps({"system": "S", "user": "U", "meta": {"k": 1}}, ensure_ascii=False)
+ "\n\n"
+ json.dumps({"system": "S", "user": "U2", "schema": {"type": "object"}})
+ "\n",
encoding="utf-8",
)
peticiones = leer_especificaciones(path)
assert len(peticiones) == 2 # la línea en blanco se ignora
assert peticiones[0].meta == {"k": 1}
assert peticiones[1].schema == {"type": "object"}
def test_especificacion_invalida_indica_la_linea(tmp_path):
path = tmp_path / "malas.jsonl"
path.write_text('{"system": "S", "user": "U"}\n{"system": "sin user"}\n')
with pytest.raises(DistillError, match=":2:"):
leer_especificaciones(path)
def test_manifiesto_faltante_da_un_error_util(trabajo):
with pytest.raises(DistillError, match="manifiesto"):
trabajo.leer_peticiones()
def test_ida_y_vuelta_del_manifiesto(trabajo):
originales = _peticiones(3)
trabajo.guardar_peticiones(originales)
assert trabajo.leer_peticiones() == originales
# --- config ----------------------------------------------------------------
def test_la_config_por_defecto_es_valida():
cfg = load_distill_config()
assert cfg.model == "claude-opus-5"
assert 0 < cfg.batch_discount <= 1
assert cfg.requests_per_batch <= 100_000
def test_rechaza_lotes_mayores_al_tope_de_la_api():
with pytest.raises(ValueError):
DistillConfig(requests_per_batch=200_000)