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:
@@ -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)
|
||||
Reference in New Issue
Block a user