diff --git a/enlace/config/__init__.py b/enlace/config/__init__.py index 54855e5..3b57d7c 100644 --- a/enlace/config/__init__.py +++ b/enlace/config/__init__.py @@ -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. """ -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 ( AgentConfig, Config, DataConfig, + DistillConfig, HardwareConfig, ModelConfig, OptimizerConfig, @@ -21,6 +27,7 @@ __all__ = [ "AgentConfig", "Config", "DataConfig", + "DistillConfig", "HardwareConfig", "ModelConfig", "OptimizerConfig", @@ -30,4 +37,5 @@ __all__ = [ "load_agent_config", "load_config", "load_config_from_argv", + "load_distill_config", ] diff --git a/enlace/config/load.py b/enlace/config/load.py index 8cdbeaa..26a9c31 100644 --- a/enlace/config/load.py +++ b/enlace/config/load.py @@ -15,7 +15,7 @@ from typing import Any 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 # 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 -def load_agent_config( - path: str | Path | None = None, overrides: list[str] | None = None -) -> AgentConfig: - """Carga la config del runtime del agente. +CONFIGS_ROOT = Path(__file__).resolve().parents[2] / "configs" - 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 - con `${oc.env:...}`, así que los secretos y las URLs de la instalación nunca - entran en un archivo commiteado. + +def _load_single(path: Path, overrides: list[str] | None, model_cls, etiqueta: str): + """Carga un YAML suelto y lo valida. No compone capas. + + 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(): 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á # definida; para pydantic eso tiene que ser ausencia, no un valor vacío. - search = data.get("search") - if isinstance(search, dict) and search.get("searxng_url") == "": - search["searxng_url"] = None + _vaciar_a_none(data) try: - return AgentConfig.model_validate(data) + return model_cls.model_validate(data) 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: diff --git a/enlace/config/schema.py b/enlace/config/schema.py index 1f36b55..71e9284 100644 --- a/enlace/config/schema.py +++ b/enlace/config/schema.py @@ -242,6 +242,40 @@ class AgentConfig(_Base): 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): """Config raíz: la composición de todas las capas.""" diff --git a/pyproject.toml b/pyproject.toml index 10b7940..9d33588 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -21,6 +21,12 @@ data = [ "fasttext-wheel>=0.9", # filtro de idioma "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 = [ "pytest>=8.0", "ruff>=0.6", diff --git a/tests/test_distill.py b/tests/test_distill.py new file mode 100644 index 0000000..ac22577 --- /dev/null +++ b/tests/test_distill.py @@ -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)